use crate::{traits::MiddlewareTrait, Middleware};
use axum::http::HeaderName;
use ayun_config::config;
use std::str::FromStr;
use tower_http::cors::CorsLayer;
pub struct Cors;
impl MiddlewareTrait for Cors {
fn handle() -> Middleware {
Box::new(|router| {
let mut cors = CorsLayer::permissive();
let config = config::<crate::config::Server>("server")?.middleware.cors;
let allow_headers = config
.allow_headers
.iter()
.map(|header| {
HeaderName::from_str(header)
.unwrap_or_else(|err| panic!("[middleware] `Cors` applied error: {}", err))
})
.collect::<Vec<_>>();
if !allow_headers.is_empty() {
cors = cors.allow_headers(allow_headers);
}
let allow_methods = config
.allow_methods
.iter()
.map(|method| {
method
.parse::<axum::http::Method>()
.unwrap_or_else(|err| panic!("[middleware] `Cors` applied error: {}", err))
})
.collect::<Vec<_>>();
if !allow_methods.is_empty() {
cors = cors.allow_methods(allow_methods);
}
let allow_origins = config
.allow_origins
.iter()
.map(|origin| {
origin
.parse::<axum::http::HeaderValue>()
.unwrap_or_else(|err| panic!("[middleware] `Cors` applied error: {}", err))
})
.collect::<Vec<_>>();
if !allow_origins.is_empty() {
cors = cors.allow_origin(allow_origins);
}
let expose_headers = config
.expose_headers
.iter()
.map(|header| {
HeaderName::from_str(header)
.unwrap_or_else(|err| panic!("[middleware] `Cors` applied error: {}", err))
})
.collect::<Vec<_>>();
if !expose_headers.is_empty() {
cors = cors.expose_headers(expose_headers);
}
let vary = config
.vary
.iter()
.map(|header| {
HeaderName::from_str(header)
.unwrap_or_else(|err| panic!("[middleware] `Cors` applied error: {}", err))
})
.collect::<Vec<_>>();
if !vary.is_empty() {
cors = cors.vary(vary);
}
Ok(router.layer(
cors.allow_credentials(config.allow_credentials)
.allow_private_network(config.allow_private_network)
.max_age(std::time::Duration::from_millis(config.max_age)),
))
})
}
}