use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct CorsConfig {
pub allowed_origins: Vec<String>,
pub allowed_methods: Vec<String>,
pub allowed_headers: Vec<String>,
}
impl CorsConfig {
pub fn validate(&self) -> Result<(), crate::config::ConfigError> {
if self.allowed_origins.is_empty() {
return Err(crate::config::ConfigError::ValidationError(
"CORS allowed_origins cannot be empty".into(),
));
}
for origin in &self.allowed_origins {
if !origin.starts_with("http://") && !origin.starts_with("https://") {
return Err(crate::config::ConfigError::ValidationError(format!(
"Invalid CORS origin: {}. Must start with http:// or https://",
origin
)));
}
let after_scheme = origin.split("://").nth(1).unwrap_or("");
if after_scheme.is_empty() {
return Err(crate::config::ConfigError::ValidationError(format!(
"Invalid CORS origin: {}. Must include host (e.g. http://example.com)",
origin
)));
}
}
Ok(())
}
}
fn parse_allowed_methods(
methods: &[String],
) -> Result<tower_http::cors::AllowMethods, crate::config::ConfigError> {
use tower_http::cors::Any;
if methods.is_empty()
|| methods
.iter()
.any(|m| m == "*" || m.eq_ignore_ascii_case("any"))
{
return Ok(Any.into());
}
let mut parsed = Vec::with_capacity(methods.len());
for m in methods {
let method = axum::http::Method::from_bytes(m.as_bytes()).map_err(|_| {
crate::config::ConfigError::ValidationError(format!(
"Invalid CORS method: {}. Use standard HTTP method names (e.g. GET, POST) or \"*\".",
m
))
})?;
parsed.push(method);
}
Ok(tower_http::cors::AllowMethods::list(parsed))
}
fn parse_allowed_headers(
headers: &[String],
) -> Result<tower_http::cors::AllowHeaders, crate::config::ConfigError> {
use tower_http::cors::Any;
if headers.is_empty() || headers.iter().any(|h| h == "*") {
return Ok(Any.into());
}
let mut parsed = Vec::with_capacity(headers.len());
for h in headers {
let name =
axum::http::HeaderName::from_bytes(h.to_lowercase().as_bytes()).map_err(|_| {
crate::config::ConfigError::ValidationError(format!(
"Invalid CORS header: {}. Use valid header names (e.g. Content-Type) or \"*\".",
h
))
})?;
parsed.push(name);
}
Ok(tower_http::cors::AllowHeaders::list(parsed))
}
pub fn build_cors_layer(
config: &CorsConfig,
) -> Result<tower_http::cors::CorsLayer, crate::config::ConfigError> {
use tower_http::cors::CorsLayer;
if config.allowed_origins.is_empty() {
return Err(crate::config::ConfigError::ValidationError(
"CORS allowed_origins cannot be empty. Use explicit origin list or disable CORS".into(),
));
}
for origin in &config.allowed_origins {
if !origin.starts_with("http://") && !origin.starts_with("https://") {
return Err(crate::config::ConfigError::ValidationError(format!(
"Invalid CORS origin: {}. Must start with http:// or https://",
origin
)));
}
let after_scheme = origin.split("://").nth(1).unwrap_or("");
if after_scheme.is_empty() {
return Err(crate::config::ConfigError::ValidationError(format!(
"Invalid CORS origin: {}. Must include host (e.g. http://example.com)",
origin
)));
}
}
let cors = CorsLayer::new()
.allow_methods(parse_allowed_methods(&config.allowed_methods)?)
.allow_headers(parse_allowed_headers(&config.allowed_headers)?);
let origins: Vec<_> = config
.allowed_origins
.iter()
.filter_map(|origin| origin.parse().ok())
.collect();
if origins.is_empty() {
return Err(crate::config::ConfigError::ValidationError(
"No valid origins found in CORS configuration".into(),
));
}
let cors = cors.allow_origin(origins);
Ok(cors)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_cors_config_with_origins() {
let json = r#"{
"allowed_origins": ["http://localhost:3000", "https://example.com"],
"allowed_methods": ["GET", "POST"],
"allowed_headers": ["Content-Type", "Authorization"]
}"#;
let config: CorsConfig = serde_json::from_str(json).unwrap();
assert_eq!(config.allowed_origins.len(), 2);
assert!(config.allowed_methods.contains(&"GET".to_string()));
assert!(config.allowed_headers.contains(&"Content-Type".to_string()));
}
#[test]
fn test_build_cors_layer_empty_origins() {
let config = CorsConfig::default();
let layer = build_cors_layer(&config);
assert!(layer.is_err());
}
#[test]
fn test_build_cors_layer_valid_origins() {
let json = r#"{"allowed_origins": ["http://localhost:3000"], "allowed_methods": [], "allowed_headers": []}"#;
let config: CorsConfig = serde_json::from_str(json).unwrap();
let layer = build_cors_layer(&config);
assert!(layer.is_ok());
}
#[test]
fn test_cors_config_validate_empty_origins() {
let config = CorsConfig {
allowed_origins: vec![],
allowed_methods: vec!["GET".to_string()],
allowed_headers: vec![],
};
let result = config.validate();
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("empty"));
}
#[test]
fn test_cors_config_validate_invalid_origin_no_scheme() {
let config = CorsConfig {
allowed_origins: vec!["localhost:3000".to_string()],
allowed_methods: vec!["GET".to_string()],
allowed_headers: vec![],
};
let result = config.validate();
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("Invalid CORS origin")
);
}
#[test]
fn test_cors_config_validate_invalid_origin_http_only() {
let config = CorsConfig {
allowed_origins: vec!["http://".to_string()],
allowed_methods: vec!["GET".to_string()],
allowed_headers: vec![],
};
let result = config.validate();
assert!(
result.is_err(),
"origin without host should be rejected, got: {:?}",
result
);
}
#[test]
fn test_cors_config_validate_valid_origins() {
let config = CorsConfig {
allowed_origins: vec![
"http://localhost:3000".to_string(),
"https://example.com".to_string(),
],
allowed_methods: vec!["GET".to_string(), "POST".to_string()],
allowed_headers: vec!["Content-Type".to_string()],
};
assert!(config.validate().is_ok());
}
#[test]
fn test_cors_config_clone() {
let config = CorsConfig {
allowed_origins: vec!["http://localhost:3000".to_string()],
allowed_methods: vec!["GET".to_string()],
allowed_headers: vec!["Authorization".to_string()],
};
let cloned = config.clone();
assert_eq!(cloned.allowed_origins, config.allowed_origins);
assert_eq!(cloned.allowed_methods, config.allowed_methods);
}
#[test]
fn test_build_cors_layer_invalid_origin_format() {
let config = CorsConfig {
allowed_origins: vec!["localhost:3000".to_string()],
allowed_methods: vec!["GET".to_string()],
allowed_headers: vec![],
};
let result = build_cors_layer(&config);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("Invalid CORS origin")
);
}
#[test]
fn test_build_cors_layer_no_valid_origins() {
let config = CorsConfig {
allowed_origins: vec!["http://\ninvalid".to_string()],
allowed_methods: vec!["GET".to_string()],
allowed_headers: vec![],
};
let result = build_cors_layer(&config);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("No valid origins"));
}
#[test]
fn test_build_cors_layer_empty_host_rejected() {
let config = CorsConfig {
allowed_origins: vec!["http://".to_string()],
allowed_methods: vec!["GET".to_string()],
allowed_headers: vec![],
};
let result = build_cors_layer(&config);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("Must include host")
);
}
#[test]
fn test_build_cors_layer_methods_config_enforced() {
let config = CorsConfig {
allowed_origins: vec!["http://localhost:3000".to_string()],
allowed_methods: vec![],
allowed_headers: vec!["Content Type".to_string()],
};
let result = build_cors_layer(&config);
assert!(result.is_err(), "invalid header must fail closed");
assert!(
result
.unwrap_err()
.to_string()
.contains("Invalid CORS header")
);
let config = CorsConfig {
allowed_origins: vec!["http://localhost:3000".to_string()],
allowed_methods: vec!["GET".to_string(), "POST".to_string()],
allowed_headers: vec!["Content-Type".to_string()],
};
assert!(build_cors_layer(&config).is_ok());
let config = CorsConfig {
allowed_origins: vec!["http://localhost:3000".to_string()],
allowed_methods: vec!["*".to_string()],
allowed_headers: vec!["*".to_string()],
};
assert!(build_cors_layer(&config).is_ok());
}
#[test]
fn test_parse_allowed_headers_rejects_invalid_name() {
let headers = vec!["Content Type".to_string()]; let result = parse_allowed_headers(&headers);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("Invalid CORS header")
);
}
}
#[cfg(all(test, feature = "http", feature = "tokio"))]
mod cors_behavior_tests {
use super::*;
use axum::http::{Method, Request};
use tower::ServiceExt;
fn cors_router(config: &CorsConfig) -> axum::Router {
use axum::Router;
use axum::routing::get;
let layer = build_cors_layer(config).expect("valid cors config");
Router::new()
.route("/x", get(|| async { "ok" }))
.layer(layer)
}
fn preflight(method: &str) -> Request<axum::body::Body> {
Request::builder()
.method(Method::OPTIONS)
.uri("/x")
.header("Origin", "http://localhost:3000")
.header("Access-Control-Request-Method", method)
.body(axum::body::Body::empty())
.unwrap()
}
#[tokio::test]
async fn preflight_rejects_method_outside_allow_list() {
let config = CorsConfig {
allowed_origins: vec!["http://localhost:3000".to_string()],
allowed_methods: vec!["GET".to_string(), "POST".to_string()],
allowed_headers: vec![],
};
let app = cors_router(&config);
let resp = app.oneshot(preflight("DELETE")).await.unwrap();
let allow_methods = resp
.headers()
.get("access-control-allow-methods")
.expect("allow-methods header must reflect explicit config")
.to_str()
.unwrap();
let echoed: Vec<&str> = allow_methods.split(',').map(str::trim).collect();
assert_eq!(
echoed,
vec!["GET", "POST"],
"config methods must be echoed exactly, got: {}",
allow_methods
);
}
#[tokio::test]
async fn preflight_allows_method_in_allow_list() {
let config = CorsConfig {
allowed_origins: vec!["http://localhost:3000".to_string()],
allowed_methods: vec!["GET".to_string(), "POST".to_string()],
allowed_headers: vec![],
};
let app = cors_router(&config);
let resp = app.oneshot(preflight("POST")).await.unwrap();
assert!(
resp.headers().get("access-control-allow-origin").is_some(),
"allowed preflight must carry allow-origin header"
);
}
}