use std::time::Duration;
use axum::http::{HeaderName, HeaderValue, Method, header};
use tower_http::cors::{AllowOrigin, CorsLayer};
pub const CORS_EXPOSE_RESPONSE_HEADERS: &[&str] =
&["x-conversation-id", "x-stream-job-id", "x-request-id"];
pub fn parse_cors_origin_header_values(allowed_origins: &[String]) -> Vec<HeaderValue> {
allowed_origins
.iter()
.map(|o| o.trim())
.filter(|o| !o.is_empty())
.filter_map(|o| HeaderValue::from_str(o).ok())
.collect()
}
pub fn try_cors_layer(allowed_origins: &[String]) -> Option<CorsLayer> {
let origins = parse_cors_origin_header_values(allowed_origins);
if origins.is_empty() {
return None;
}
let expose: Vec<HeaderName> = CORS_EXPOSE_RESPONSE_HEADERS
.iter()
.map(|n| HeaderName::from_static(n))
.collect();
Some(
CorsLayer::new()
.allow_origin(AllowOrigin::list(origins))
.allow_methods([
Method::GET,
Method::POST,
Method::PUT,
Method::DELETE,
Method::OPTIONS,
])
.allow_headers([
header::AUTHORIZATION,
header::CONTENT_TYPE,
header::ACCEPT,
HeaderName::from_static("x-api-key"),
HeaderName::from_static("last-event-id"),
HeaderName::from_static("x-crabmate-github-token"),
HeaderName::from_static("x-crabmate-github-token-delivery"),
])
.expose_headers(expose)
.max_age(Duration::from_secs(600))
.allow_credentials(true),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_allowlist_yields_no_layer() {
assert!(try_cors_layer(&[]).is_none());
assert!(try_cors_layer(&[String::new(), " ".into()]).is_none());
}
#[test]
fn parses_and_trims_origins() {
let vals = parse_cors_origin_header_values(&[
" http://127.0.0.1:8081 ".into(),
"".into(),
"https://ui.example.com".into(),
" ".into(),
]);
assert_eq!(vals.len(), 2);
assert_eq!(vals[0].to_str().unwrap(), "http://127.0.0.1:8081");
assert_eq!(vals[1].to_str().unwrap(), "https://ui.example.com");
assert!(try_cors_layer(&["http://127.0.0.1:8081".into()]).is_some());
}
#[test]
fn expose_list_covers_stream_session_headers() {
assert!(CORS_EXPOSE_RESPONSE_HEADERS.contains(&"x-conversation-id"));
assert!(CORS_EXPOSE_RESPONSE_HEADERS.contains(&"x-stream-job-id"));
assert!(CORS_EXPOSE_RESPONSE_HEADERS.contains(&"x-request-id"));
for name in CORS_EXPOSE_RESPONSE_HEADERS {
assert!(
HeaderName::from_static(name).as_str() == *name,
"invalid expose header name: {name}"
);
}
}
}