Skip to main content

rskit_http/
cors.rs

1//! CORS policy and Tower layer construction.
2
3use std::time::Duration;
4
5use http::{HeaderName, HeaderValue, Method};
6use rskit_errors::{AppError, AppResult};
7use tower_http::cors::{AllowHeaders, AllowMethods, AllowOrigin, CorsLayer};
8
9/// Cross-origin resource sharing policy.
10///
11/// The default is deny-by-default: no origins, methods, headers,
12/// or credentials are allowed unless explicitly configured.
13#[derive(Debug, Clone, Default, serde::Deserialize)]
14#[serde(default)]
15pub struct CorsPolicy {
16    /// Allowed Origin header values.
17    pub allowed_origins: Vec<String>,
18    /// Allowed HTTP methods.
19    pub allowed_methods: Vec<String>,
20    /// Allowed request headers.
21    pub allowed_headers: Vec<String>,
22    /// Whether to allow credentials.
23    pub allow_credentials: bool,
24    /// Cache duration for pre-flight responses.
25    pub max_age: Duration,
26}
27
28impl CorsPolicy {
29    /// Validate the CORS policy.
30    ///
31    /// # Errors
32    /// Returns an error when an origin, method, or header is invalid.
33    pub fn validate(&self) -> AppResult<()> {
34        self.build_origins()?;
35        self.build_methods()?;
36        self.build_headers()?;
37        Ok(())
38    }
39
40    /// Build a Tower CORS layer from this policy.
41    ///
42    /// # Errors
43    /// Returns an error when an origin, method, or header is invalid.
44    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}