use tower_http::cors::CorsLayer;
use crate::config::{ApiGatewayConfig, CorsConfig};
pub fn validate_cors_config(cfg: &ApiGatewayConfig) -> Result<(), String> {
if !cfg.cors_enabled {
return Ok(());
}
let cors_cfg: CorsConfig = cfg.cors.clone().unwrap_or_default();
if !cors_cfg.allow_credentials {
return Ok(());
}
let wildcard_lists = [
("cors.allowed_origins", &cors_cfg.allowed_origins),
("cors.allowed_methods", &cors_cfg.allowed_methods),
("cors.allowed_headers", &cors_cfg.allowed_headers),
("cors.exposed_headers", &cors_cfg.exposed_headers),
];
for (name, list) in wildcard_lists {
if list.iter().any(|v| v == "*") {
return Err(format!(
"invalid CORS configuration: `{name}` contains the wildcard \"*\" while \
`cors.allow_credentials` is true; browsers forbid this combination and \
tower-http would panic at startup — list the values explicitly or \
disable credentials"
));
}
}
Ok(())
}
fn parse_list<T: std::str::FromStr>(list_name: &'static str, items: &[String]) -> Vec<T> {
items
.iter()
.filter_map(|s| {
s.parse::<T>()
.map_err(|_| {
tracing::warn!(
entry = %s,
list = list_name,
"api-gateway CORS config entry does not parse; ignoring it"
);
})
.ok()
})
.collect()
}
pub fn build_cors_layer(cfg: &ApiGatewayConfig) -> CorsLayer {
let cors_cfg: CorsConfig = cfg.cors.clone().unwrap_or_default();
let mut layer = CorsLayer::new();
if cors_cfg.allowed_origins.iter().any(|o| o == "*") {
layer = layer.allow_origin(tower_http::cors::Any);
} else {
let origins: Vec<axum::http::HeaderValue> =
parse_list("allowed_origins", &cors_cfg.allowed_origins);
if !origins.is_empty() {
layer = layer.allow_origin(origins);
}
}
if cors_cfg.allowed_methods.iter().any(|m| m == "*") {
layer = layer.allow_methods(tower_http::cors::Any);
} else {
let methods: Vec<axum::http::Method> =
parse_list("allowed_methods", &cors_cfg.allowed_methods);
if !methods.is_empty() {
layer = layer.allow_methods(methods);
}
}
if cors_cfg.allowed_headers.iter().any(|h| h == "*") {
layer = layer.allow_headers(tower_http::cors::Any);
} else {
let headers: Vec<axum::http::HeaderName> =
parse_list("allowed_headers", &cors_cfg.allowed_headers);
if !headers.is_empty() {
layer = layer.allow_headers(headers);
}
}
if cors_cfg.exposed_headers.iter().any(|h| h == "*") {
layer = layer.expose_headers(tower_http::cors::Any);
} else {
let headers: Vec<axum::http::HeaderName> =
parse_list("exposed_headers", &cors_cfg.exposed_headers);
if !headers.is_empty() {
layer = layer.expose_headers(headers);
}
}
if cors_cfg.allow_credentials {
layer = layer.allow_credentials(true);
}
if cors_cfg.max_age_seconds > 0 {
layer = layer.max_age(std::time::Duration::from_secs(cors_cfg.max_age_seconds));
}
layer
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::ApiGatewayConfig;
fn cfg_with(cors: CorsConfig) -> ApiGatewayConfig {
ApiGatewayConfig {
cors_enabled: true,
cors: Some(cors),
..Default::default()
}
}
#[test]
fn wildcard_exposed_headers_with_credentials_rejected_at_validation() {
let cfg = cfg_with(CorsConfig {
allowed_origins: vec!["https://ui.example.com".to_owned()],
allowed_methods: vec!["GET".to_owned()],
allowed_headers: vec!["Content-Type".to_owned()],
exposed_headers: vec!["*".to_owned()],
allow_credentials: true,
..Default::default()
});
let err = validate_cors_config(&cfg).expect_err("wildcard+credentials must be rejected");
assert!(err.contains("exposed_headers"), "got: {err}");
}
#[test]
fn wildcard_origins_with_credentials_rejected_at_validation() {
let cfg = cfg_with(CorsConfig {
allow_credentials: true,
..Default::default() });
assert!(validate_cors_config(&cfg).is_err());
}
#[test]
fn explicit_lists_with_credentials_pass_validation() {
let cfg = cfg_with(CorsConfig {
allowed_origins: vec!["https://ui.example.com".to_owned()],
allowed_methods: vec!["GET".to_owned(), "POST".to_owned()],
allowed_headers: vec!["Content-Type".to_owned()],
exposed_headers: vec!["ETag".to_owned()],
allow_credentials: true,
..Default::default()
});
assert!(validate_cors_config(&cfg).is_ok());
}
#[test]
fn cors_disabled_skips_validation() {
let mut cfg = cfg_with(CorsConfig {
exposed_headers: vec!["*".to_owned()],
allow_credentials: true,
..Default::default()
});
cfg.cors_enabled = false;
assert!(validate_cors_config(&cfg).is_ok());
}
}