use crate::{DisplayConfig, SwordLayerRegistrar};
use axum::http::{HeaderName, HeaderValue, Method};
use serde::{Deserialize, Serialize};
use thisconfig::Config;
use thisconfig::{ConfigItem, TimeConfig};
use tower_http::cors::Any;
pub use tower_http::cors::CorsLayer;
#[derive(Clone, Debug, Deserialize, Serialize)]
#[serde(default)]
pub struct CorsConfig {
#[serde(rename = "allow-origins")]
pub allow_origins: Option<Vec<String>>,
#[serde(rename = "allow-methods")]
pub allow_methods: Option<Vec<String>>,
#[serde(rename = "allow-headers")]
pub allow_headers: Option<Vec<String>>,
#[serde(rename = "allow-credentials")]
pub allow_credentials: Option<bool>,
#[serde(rename = "max-age")]
pub max_age: Option<TimeConfig>,
pub display: bool,
}
impl Default for CorsConfig {
fn default() -> Self {
Self {
allow_origins: Some(vec!["*".into()]),
allow_methods: Some(vec![
"GET".into(),
"POST".into(),
"PUT".into(),
"DELETE".into(),
"PATCH".into(),
"OPTIONS".into(),
"HEAD".into(),
]),
allow_headers: Some(vec!["*".into()]),
allow_credentials: None,
max_age: None,
display: false,
}
}
}
impl DisplayConfig for CorsConfig {
fn display(&self) {
if !self.display {
return;
}
tracing::debug!(
target: "sword.layers.cors",
allow_origins = ?self.allow_origins,
allow_methods = ?self.allow_methods,
allow_headers = ?self.allow_headers,
allow_credentials = self.allow_credentials,
max_age = ?self.max_age.as_ref().map(|value| &value.raw),
"CORS layer configuration"
);
}
}
impl From<CorsConfig> for CorsLayer {
fn from(config: CorsConfig) -> CorsLayer {
let mut layer = CorsLayer::new();
if let Some(allow_credentials) = config.allow_credentials {
layer = layer.allow_credentials(allow_credentials);
}
if let Some(origin) = &config.allow_origins {
if origin.iter().any(|o| o == "*") {
layer = layer.allow_origin(Any);
} else {
let parsed_origin: Vec<HeaderValue> = origin
.iter()
.filter_map(|o| HeaderValue::from_str(o).ok())
.collect();
layer = layer.allow_origin(parsed_origin);
}
}
if let Some(methods) = &config.allow_methods {
if methods.iter().any(|m| m == "*") {
layer = layer.allow_methods(Any);
} else {
let parsed_methods: Vec<Method> =
methods.iter().filter_map(|m| m.parse().ok()).collect();
layer = layer.allow_methods(parsed_methods);
}
}
if let Some(headers) = &config.allow_headers {
if headers.iter().any(|h| h == "*") {
layer = layer.allow_headers(Any);
} else {
let parsed_headers: Vec<HeaderName> =
headers.iter().filter_map(|h| h.parse().ok()).collect();
layer = layer.allow_headers(parsed_headers);
}
}
if let Some(max_age) = &config.max_age {
layer = layer.max_age(max_age.parsed);
}
layer
}
}
impl ConfigItem for CorsConfig {
fn key() -> &'static str {
"cors"
}
}
inventory::submit! {
SwordLayerRegistrar {
name: "cors",
register: |config: &Config| {
let layer: CorsLayer = config.get_or_default::<CorsConfig>().into();
Box::new(move |any: &mut dyn std::any::Any| {
let stack = any
.downcast_mut::<sword_core::LayerStack<sword_core::State>>()
.expect("SwordLayerRegistrar: expected LayerStack<State>");
stack.push(layer);
})
},
display: |config: &Config| {
config.get_or_default::<CorsConfig>().display();
},
}
}