1use std::time::Duration;
4
5use http::{HeaderName, HeaderValue, Method};
6use rskit_errors::{AppError, AppResult};
7use tower_http::cors::{AllowHeaders, AllowMethods, AllowOrigin, CorsLayer};
8
9#[derive(Debug, Clone, Default, serde::Deserialize)]
14#[serde(default)]
15pub struct CorsPolicy {
16 pub allowed_origins: Vec<String>,
18 pub allowed_methods: Vec<String>,
20 pub allowed_headers: Vec<String>,
22 pub allow_credentials: bool,
24 pub max_age: Duration,
26}
27
28impl CorsPolicy {
29 pub fn validate(&self) -> AppResult<()> {
34 self.build_origins()?;
35 self.build_methods()?;
36 self.build_headers()?;
37 Ok(())
38 }
39
40 pub fn layer(&self) -> AppResult<CorsLayer> {
45 let origins = self.build_origins()?;
46 let methods = self.build_methods()?;
47 let headers = self.build_headers()?;
48
49 let mut layer = CorsLayer::new()
50 .allow_origin(AllowOrigin::list(origins))
51 .allow_methods(AllowMethods::list(methods))
52 .allow_headers(AllowHeaders::list(headers))
53 .allow_credentials(self.allow_credentials);
54
55 if !self.max_age.is_zero() {
56 layer = layer.max_age(self.max_age);
57 }
58
59 Ok(layer)
60 }
61
62 fn build_origins(&self) -> AppResult<Vec<HeaderValue>> {
63 if self.allow_credentials && self.allowed_origins.is_empty() {
64 return Err(AppError::invalid_input(
65 "allowed_origins",
66 "credentials require an explicit origin allow-list",
67 ));
68 }
69 self.allowed_origins
70 .iter()
71 .map(|origin| {
72 validate_allowed_origin(origin)?;
73 HeaderValue::from_str(origin).map_err(|error| {
74 AppError::invalid_input("allowed_origins", format!("invalid origin: {error}"))
75 })
76 })
77 .collect()
78 }
79
80 fn build_methods(&self) -> AppResult<Vec<Method>> {
81 self.allowed_methods
82 .iter()
83 .map(|method| {
84 method.parse::<Method>().map_err(|error| {
85 AppError::invalid_input("allowed_methods", format!("invalid method: {error}"))
86 })
87 })
88 .collect()
89 }
90
91 fn build_headers(&self) -> AppResult<Vec<HeaderName>> {
92 self.allowed_headers
93 .iter()
94 .map(|header| {
95 HeaderName::from_bytes(header.as_bytes()).map_err(|error| {
96 AppError::invalid_input("allowed_headers", format!("invalid header: {error}"))
97 })
98 })
99 .collect()
100 }
101}
102
103fn validate_allowed_origin(origin: &str) -> AppResult<()> {
104 if origin == "*" {
105 return Err(AppError::invalid_input(
106 "allowed_origins",
107 "wildcard origins are not allowed",
108 ));
109 }
110
111 let parsed = origin.parse::<http::Uri>().map_err(|error| {
112 AppError::invalid_input("allowed_origins", format!("invalid origin: {error}"))
113 })?;
114
115 match parsed.scheme_str() {
116 Some("http" | "https") => {}
117 scheme => {
118 return Err(AppError::invalid_input(
119 "allowed_origins",
120 format!("origin scheme must be http or https, got {scheme:?}"),
121 ));
122 }
123 }
124
125 let Some(authority) = parsed.authority() else {
126 return Err(AppError::invalid_input(
127 "allowed_origins",
128 "origin must include a host",
129 ));
130 };
131 if authority.as_str().contains('@') {
132 return Err(AppError::invalid_input(
133 "allowed_origins",
134 "origin must not contain credentials",
135 ));
136 }
137
138 let path = parsed.path();
139 if path != "/" && !path.is_empty() {
140 return Err(AppError::invalid_input(
141 "allowed_origins",
142 "origin must not contain a path",
143 ));
144 }
145 if parsed.query().is_some() {
146 return Err(AppError::invalid_input(
147 "allowed_origins",
148 "origin must not contain a query",
149 ));
150 }
151 if origin.contains('#') {
152 return Err(AppError::invalid_input(
153 "allowed_origins",
154 "origin must not contain a fragment",
155 ));
156 }
157 Ok(())
158}
159
160#[cfg(test)]
161mod tests {
162 use super::*;
163
164 #[test]
165 fn defaults_are_deny_by_default() {
166 let policy = CorsPolicy::default();
167 assert!(policy.allowed_origins.is_empty());
168 assert!(policy.allowed_methods.is_empty());
169 assert!(policy.allowed_headers.is_empty());
170 assert!(policy.validate().is_ok());
171 }
172
173 #[test]
174 fn omitted_fields_deserialize_to_deny_by_default_values() {
175 let policy: CorsPolicy =
176 serde_json::from_str(r#"{"allowed_origins":["https://example.com"]}"#).unwrap();
177 assert_eq!(policy.allowed_origins, vec!["https://example.com"]);
178 assert!(policy.allowed_methods.is_empty());
179 assert!(policy.allowed_headers.is_empty());
180 assert!(!policy.allow_credentials);
181 assert!(policy.max_age.is_zero());
182 }
183
184 #[test]
185 fn rejects_wildcard_origin_before_url_validation() {
186 let policy = CorsPolicy {
187 allowed_origins: vec!["*".to_string()],
188 allowed_methods: vec!["GET".to_string()],
189 allowed_headers: vec!["authorization".to_string()],
190 allow_credentials: false,
191 max_age: Duration::from_mins(1),
192 };
193 let err = policy.validate().unwrap_err();
194 assert!(err.to_string().contains("wildcard origins are not allowed"));
195 }
196
197 #[test]
198 fn rejects_invalid_methods_and_headers() {
199 let policy = CorsPolicy {
200 allowed_origins: vec!["https://example.com".to_string()],
201 allowed_methods: vec!["bad method".to_string()],
202 allowed_headers: vec!["authorization".to_string()],
203 allow_credentials: false,
204 max_age: Duration::from_mins(1),
205 };
206 assert!(policy.validate().is_err());
207
208 let policy = CorsPolicy {
209 allowed_origins: vec!["https://example.com".to_string()],
210 allowed_methods: vec!["GET".to_string()],
211 allowed_headers: vec!["bad header".to_string()],
212 allow_credentials: false,
213 max_age: Duration::from_mins(1),
214 };
215 assert!(policy.validate().is_err());
216 }
217
218 #[test]
219 fn rejects_malformed_and_hostless_origins() {
220 for origin in ["http://[", "https:"] {
221 let policy = CorsPolicy {
222 allowed_origins: vec![origin.to_string()],
223 allowed_methods: vec!["GET".to_string()],
224 allowed_headers: vec!["authorization".to_string()],
225 allow_credentials: false,
226 max_age: Duration::from_secs(0),
227 };
228 assert!(policy.validate().is_err());
229 }
230 }
231
232 #[test]
233 fn layer_builds_with_and_without_max_age() {
234 let policy = CorsPolicy {
235 allowed_origins: vec!["https://example.com".to_string()],
236 allowed_methods: vec!["GET".to_string(), "POST".to_string()],
237 allowed_headers: vec!["authorization".to_string()],
238 allow_credentials: true,
239 max_age: Duration::from_mins(5),
240 };
241 assert!(policy.layer().is_ok());
242
243 let policy = CorsPolicy {
244 max_age: Duration::from_secs(0),
245 ..policy
246 };
247 assert!(policy.layer().is_ok());
248 }
249
250 #[test]
251 fn layer_surfaces_invalid_method() {
252 let policy = CorsPolicy {
253 allowed_origins: vec!["https://example.com".to_string()],
254 allowed_methods: vec!["bad method".to_string()],
255 allowed_headers: vec![],
256 allow_credentials: false,
257 max_age: Duration::from_secs(0),
258 };
259 assert!(policy.layer().is_err());
260 }
261
262 #[test]
263 fn credentials_require_explicit_origin() {
264 let policy = CorsPolicy {
265 allowed_origins: vec![],
266 allowed_methods: vec!["GET".to_string()],
267 allowed_headers: vec![],
268 allow_credentials: true,
269 max_age: Duration::from_secs(0),
270 };
271 let err = policy.validate().unwrap_err();
272 assert!(
273 err.to_string()
274 .contains("credentials require an explicit origin allow-list")
275 );
276 }
277}