use crate::auth::{TokenAuth, auth_middleware};
use crate::health::simple_health_handler;
use crate::{Result, Sidecar};
use axum::{
Router,
body::Body,
extract::DefaultBodyLimit,
http::{Method, Request, Response, StatusCode, header::HeaderName},
middleware::{self, Next},
routing::get,
};
use std::net::SocketAddr;
use std::time::Duration;
use tower_http::{
compression::CompressionLayer,
cors::{AllowHeaders, AllowMethods, AllowOrigin, CorsLayer},
limit::RequestBodyLimitLayer,
timeout::TimeoutLayer,
trace::TraceLayer,
};
use tracing::{info, warn};
use uuid::Uuid;
#[derive(Clone, Debug)]
pub struct TraceId(pub String);
async fn trace_id_middleware(req: Request<Body>, next: Next) -> Response<Body> {
let trace_id = req
.headers()
.get("x-trace-id")
.and_then(|v| v.to_str().ok())
.map_or_else(|| Uuid::new_v4().to_string(), String::from);
let span = tracing::info_span!(
"request",
trace_id = %trace_id,
method = %req.method(),
path = %req.uri().path()
);
let _guard = span.enter();
let mut req = req;
req.extensions_mut().insert(TraceId(trace_id.clone()));
let mut response = next.run(req).await;
if let Ok(header_value) = trace_id.parse() {
response.headers_mut().insert("x-trace-id", header_value);
}
response
}
const DEFAULT_MAX_BODY_SIZE: usize = 10 * 1024 * 1024;
#[derive(Clone, Debug, Default)]
pub struct CorsConfig {
enabled: bool,
origins: Vec<String>,
methods: Vec<Method>,
headers: Vec<HeaderName>,
}
impl CorsConfig {
#[must_use]
pub fn disabled() -> Self {
Self::default()
}
#[must_use]
pub fn permissive() -> Self {
Self {
enabled: true,
origins: Vec::new(),
methods: Vec::new(),
headers: Vec::new(),
}
}
#[must_use]
pub fn with_origins<I, S>(origins: I) -> Self
where
I: IntoIterator<Item = S>,
S: Into<String>,
{
Self {
enabled: true,
origins: origins.into_iter().map(Into::into).collect(),
methods: Vec::new(),
headers: Vec::new(),
}
}
#[must_use]
pub fn methods<I>(mut self, methods: I) -> Self
where
I: IntoIterator<Item = Method>,
{
self.methods = methods.into_iter().collect();
self
}
#[must_use]
pub fn headers<I>(mut self, headers: I) -> Self
where
I: IntoIterator<Item = HeaderName>,
{
self.headers = headers.into_iter().collect();
self
}
fn into_layer(self) -> Option<CorsLayer> {
if !self.enabled {
return None;
}
let mut layer = CorsLayer::new();
if self.origins.is_empty() {
warn!("CORS configured with allow_any_origin - this is insecure for production");
layer = layer.allow_origin(AllowOrigin::any());
} else {
let origins: Vec<_> = self.origins.iter().filter_map(|o| o.parse().ok()).collect();
layer = layer.allow_origin(origins);
}
if self.methods.is_empty() {
layer = layer.allow_methods(AllowMethods::any());
} else {
layer = layer.allow_methods(self.methods);
}
if self.headers.is_empty() {
layer = layer.allow_headers(AllowHeaders::any());
} else {
layer = layer.allow_headers(self.headers);
}
Some(layer)
}
}
#[derive(Debug)]
pub struct SidecarBuilder {
port: u16,
auth: TokenAuth,
cors_config: CorsConfig,
enable_compression: bool,
timeout_secs: u64,
max_body_size: usize,
}
impl Default for SidecarBuilder {
fn default() -> Self {
Self::new()
}
}
impl SidecarBuilder {
#[must_use]
pub fn new() -> Self {
Self {
port: 3001,
auth: TokenAuth::from_env(),
cors_config: CorsConfig::disabled(),
enable_compression: true,
timeout_secs: 30,
max_body_size: DEFAULT_MAX_BODY_SIZE,
}
}
#[must_use]
pub fn port(mut self, port: u16) -> Self {
self.port = port;
self
}
#[must_use]
pub fn auth_token(mut self, token: impl Into<String>) -> Self {
self.auth = TokenAuth::new(token);
self
}
#[must_use]
pub fn auth_token_from_env(mut self) -> Self {
self.auth = TokenAuth::from_env();
self
}
#[must_use]
pub fn no_auth(mut self) -> Self {
self.auth = TokenAuth::disabled();
self
}
#[must_use]
pub fn cors(mut self, config: CorsConfig) -> Self {
self.cors_config = config;
self
}
#[must_use]
pub fn cors_permissive(mut self) -> Self {
self.cors_config = CorsConfig::permissive();
self
}
#[must_use]
pub fn cors_disabled(mut self) -> Self {
self.cors_config = CorsConfig::disabled();
self
}
#[must_use]
pub fn compression(mut self, enabled: bool) -> Self {
self.enable_compression = enabled;
self
}
#[must_use]
pub fn timeout(mut self, secs: u64) -> Self {
self.timeout_secs = secs;
self
}
#[must_use]
pub fn max_body_size(mut self, bytes: usize) -> Self {
self.max_body_size = bytes;
self
}
#[must_use]
pub fn max_body_size_mb(self, mb: usize) -> Self {
self.max_body_size(mb * 1024 * 1024)
}
pub async fn serve<S: Sidecar>(self, sidecar: S) -> Result<()> {
let service_name = sidecar.name();
let mut app = sidecar
.router()
.route("/health", get(simple_health_handler));
let auth = self.auth.clone();
app = app.layer(middleware::from_fn(move |req, next| {
let auth = auth.clone();
auth_middleware(auth, req, next)
}));
if self.enable_compression {
app = app.layer(CompressionLayer::new());
}
let cors_enabled = self.cors_config.enabled;
if let Some(cors_layer) = self.cors_config.clone().into_layer() {
app = app.layer(cors_layer);
}
app = app
.layer(DefaultBodyLimit::max(self.max_body_size))
.layer(RequestBodyLimitLayer::new(self.max_body_size))
.layer(TimeoutLayer::with_status_code(
StatusCode::REQUEST_TIMEOUT,
Duration::from_secs(self.timeout_secs),
))
.layer(TraceLayer::new_for_http())
.layer(middleware::from_fn(trace_id_middleware));
let addr = SocketAddr::from(([0, 0, 0, 0], self.port));
let listener = tokio::net::TcpListener::bind(addr).await?;
info!(
service = service_name,
port = self.port,
auth = self.auth.is_enabled(),
cors = cors_enabled,
max_body_size = self.max_body_size,
"Starting sidecar"
);
axum::serve(listener, app).await?;
Ok(())
}
pub async fn serve_multi(self, sidecars: Vec<(&str, Box<dyn Sidecar>)>) -> Result<()> {
let service_count = sidecars.len();
let mut app = Router::new().route("/health", get(simple_health_handler));
for (prefix, sidecar) in sidecars {
info!(
service = sidecar.name(),
prefix = prefix,
"Mounting service"
);
app = app.nest(prefix, sidecar.router());
}
let auth = self.auth.clone();
app = app.layer(middleware::from_fn(move |req, next| {
let auth = auth.clone();
auth_middleware(auth, req, next)
}));
if self.enable_compression {
app = app.layer(CompressionLayer::new());
}
let cors_enabled = self.cors_config.enabled;
if let Some(cors_layer) = self.cors_config.into_layer() {
app = app.layer(cors_layer);
}
app = app
.layer(DefaultBodyLimit::max(self.max_body_size))
.layer(RequestBodyLimitLayer::new(self.max_body_size))
.layer(TimeoutLayer::with_status_code(
StatusCode::REQUEST_TIMEOUT,
Duration::from_secs(self.timeout_secs),
))
.layer(TraceLayer::new_for_http())
.layer(middleware::from_fn(trace_id_middleware));
let addr = SocketAddr::from(([0, 0, 0, 0], self.port));
let listener = tokio::net::TcpListener::bind(addr).await?;
info!(
port = self.port,
services = service_count,
cors = cors_enabled,
max_body_size = self.max_body_size,
"Starting supercar (multi-service mode)"
);
axum::serve(listener, app).await?;
Ok(())
}
}