diff --git a/crates/browser/src/loader.rs b/crates/browser/src/loader.rs index fdd8242..d653d02 100644 --- a/crates/browser/src/loader.rs +++ b/crates/browser/src/loader.rs @@ -9,7 +9,11 @@ use std::fmt; use we_encoding::sniff::sniff_encoding; use we_encoding::Encoding; use we_net::client::{ClientError, HttpClient}; -use we_net::http::ContentType; +use we_net::cors::{ + self, build_preflight_headers, check_cors_response, needs_preflight, validate_preflight, + CredentialsMode, PreflightCache, +}; +use we_net::http::{ContentType, Headers, Method}; use we_url::data_url::{is_data_url, parse_data_url}; use we_url::{Origin, Url}; @@ -109,6 +113,7 @@ pub enum Resource { /// Loads resources over HTTP/HTTPS with encoding detection and content-type handling. pub struct ResourceLoader { client: HttpClient, + preflight_cache: PreflightCache, } impl ResourceLoader { @@ -116,6 +121,7 @@ impl ResourceLoader { pub fn new() -> Self { Self { client: HttpClient::new(), + preflight_cache: PreflightCache::new(), } } @@ -205,16 +211,36 @@ impl ResourceLoader { } } - /// Fetch a subresource with Same-Origin Policy enforcement. + /// Fetch a subresource with Same-Origin Policy and CORS enforcement. /// /// Checks the resource URL's origin against the document origin. For - /// cross-origin requests without CORS headers, scripts, stylesheets, - /// fetch, and font loads are blocked. Images and navigations are allowed. + /// cross-origin requests, performs CORS checks including preflight for + /// non-simple requests. Images and navigations are always allowed. pub fn fetch_subresource( &mut self, url: &Url, document_origin: &Origin, request_type: ResourceRequestType, + ) -> Result { + self.fetch_subresource_with_cors( + url, + document_origin, + request_type, + Method::Get, + &Headers::new(), + CredentialsMode::SameOrigin, + ) + } + + /// Fetch a subresource with full CORS control (method, headers, credentials). + pub fn fetch_subresource_with_cors( + &mut self, + url: &Url, + document_origin: &Origin, + request_type: ResourceRequestType, + method: Method, + extra_headers: &Headers, + credentials_mode: CredentialsMode, ) -> Result { // data: and about: URLs are always allowed (local, no network). if url.scheme() == "data" || url.scheme() == "about" { @@ -230,11 +256,56 @@ impl ResourceLoader { let resource_origin = url.origin(); if document_origin.same_origin(&resource_origin) { - return self.fetch(url); + // Same-origin: no CORS needed, use provided method and headers. + let response = self.client.request(method, url, extra_headers, None)?; + if response.status_code >= 400 { + return Err(LoadError::HttpStatus { + status: response.status_code, + reason: response.reason.clone(), + }); + } + return decode_response(response, url); } - // Cross-origin: perform the fetch but check for CORS headers. - let response = self.client.get(url)?; + // Cross-origin: CORS flow. + let url_str = url.serialize(); + + // Add Origin header to the request. + let mut request_headers = Headers::new(); + for (name, value) in extra_headers.iter() { + request_headers.add(name, value); + } + request_headers.set("Origin", &document_origin.serialize()); + + // Check if a preflight is needed. + if needs_preflight(method, extra_headers) { + // Check the preflight cache first. + if !self + .preflight_cache + .lookup(document_origin, &url_str, method, extra_headers) + { + // Send preflight OPTIONS request. + let pf_headers = build_preflight_headers(document_origin, method, extra_headers); + let pf_response = self + .client + .request(Method::Options, url, &pf_headers, None)?; + + let pf_result = + validate_preflight(&pf_response.headers, document_origin, credentials_mode) + .map_err(|reason| LoadError::CrossOriginBlocked { url: reason })?; + + // Check that the actual request's method/headers are allowed. + cors::preflight_allows(&pf_result, method, extra_headers) + .map_err(|reason| LoadError::CrossOriginBlocked { url: reason })?; + + // Cache the preflight result. + self.preflight_cache + .store(document_origin, &url_str, &pf_result); + } + } + + // Perform the actual request. + let response = self.client.request(method, url, &request_headers, None)?; if response.status_code >= 400 { return Err(LoadError::HttpStatus { @@ -243,79 +314,11 @@ impl ResourceLoader { }); } - // Check Access-Control-Allow-Origin header. - let doc_origin_str = document_origin.serialize(); - let allowed = response - .headers - .get("access-control-allow-origin") - .map(|v| { - let v = v.trim(); - v == "*" || v == doc_origin_str - }) - .unwrap_or(false); - - if !allowed { - return Err(LoadError::CrossOriginBlocked { - url: url.serialize(), - }); - } - - // CORS allows it — decode as normal. - let content_type = response.content_type(); - let mime = content_type - .as_ref() - .map(|ct| ct.mime_type.as_str()) - .unwrap_or("application/octet-stream"); + // Check CORS response headers. + check_cors_response(&response.headers, document_origin, credentials_mode) + .map_err(|reason| LoadError::CrossOriginBlocked { url: reason })?; - match classify_mime(mime) { - MimeClass::Html => { - let (text, encoding) = - decode_text_resource(&response.body, content_type.as_ref(), true); - Ok(Resource::Html { - text, - base_url: url.clone(), - encoding, - }) - } - MimeClass::Css => { - let (text, _encoding) = - decode_text_resource(&response.body, content_type.as_ref(), false); - Ok(Resource::Css { - text, - url: url.clone(), - }) - } - MimeClass::Script => { - let (text, _encoding) = - decode_text_resource(&response.body, content_type.as_ref(), false); - Ok(Resource::Script { - text, - url: url.clone(), - }) - } - MimeClass::Image => Ok(Resource::Image { - data: response.body, - mime_type: mime.to_string(), - url: url.clone(), - }), - MimeClass::Other => { - if mime.starts_with("text/") { - let (text, _encoding) = - decode_text_resource(&response.body, content_type.as_ref(), false); - Ok(Resource::Other { - data: text.into_bytes(), - mime_type: mime.to_string(), - url: url.clone(), - }) - } else { - Ok(Resource::Other { - data: response.body, - mime_type: mime.to_string(), - url: url.clone(), - }) - } - } - } + decode_response(response, url) } /// Fetch a URL string, resolving it against an optional base URL. @@ -380,6 +383,65 @@ impl Default for ResourceLoader { } } +/// Decode an HTTP response into a Resource based on its Content-Type. +fn decode_response(response: we_net::http::HttpResponse, url: &Url) -> Result { + let content_type = response.content_type(); + let mime = content_type + .as_ref() + .map(|ct| ct.mime_type.as_str()) + .unwrap_or("application/octet-stream"); + + match classify_mime(mime) { + MimeClass::Html => { + let (text, encoding) = + decode_text_resource(&response.body, content_type.as_ref(), true); + Ok(Resource::Html { + text, + base_url: url.clone(), + encoding, + }) + } + MimeClass::Css => { + let (text, _encoding) = + decode_text_resource(&response.body, content_type.as_ref(), false); + Ok(Resource::Css { + text, + url: url.clone(), + }) + } + MimeClass::Script => { + let (text, _encoding) = + decode_text_resource(&response.body, content_type.as_ref(), false); + Ok(Resource::Script { + text, + url: url.clone(), + }) + } + MimeClass::Image => Ok(Resource::Image { + data: response.body, + mime_type: mime.to_string(), + url: url.clone(), + }), + MimeClass::Other => { + if mime.starts_with("text/") { + let (text, _encoding) = + decode_text_resource(&response.body, content_type.as_ref(), false); + Ok(Resource::Other { + data: text.into_bytes(), + mime_type: mime.to_string(), + url: url.clone(), + }) + } else { + Ok(Resource::Other { + data: response.body, + mime_type: mime.to_string(), + url: url.clone(), + }) + } + } + } +} + // --------------------------------------------------------------------------- // MIME classification // --------------------------------------------------------------------------- diff --git a/crates/js/src/fetch.rs b/crates/js/src/fetch.rs index 0584bbd..29f6f71 100644 --- a/crates/js/src/fetch.rs +++ b/crates/js/src/fetch.rs @@ -117,10 +117,12 @@ pub fn fetch_native(args: &[Value], ctx: &mut NativeContext) -> Result = Vec::new(); let mut body: Option> = None; + let mut cors_mode = "cors".to_string(); + let mut credentials_mode = "same-origin".to_string(); if let Some(Value::Object(opts_ref)) = args.get(1) { if let Some(HeapObject::Object(data)) = ctx.gc.get(*opts_ref) { @@ -150,6 +152,18 @@ pub fn fetch_native(args: &[Value], ctx: &mut NativeContext) -> Result Result Result we_net::cors::CredentialsMode { + match mode { + "include" => we_net::cors::CredentialsMode::Include, + "omit" => we_net::cors::CredentialsMode::Omit, + _ => we_net::cors::CredentialsMode::SameOrigin, + } +} + /// Perform the actual HTTP fetch (runs on a background thread). /// -/// If `document_origin` is set, performs Same-Origin Policy checks. Cross-origin -/// responses without a matching `Access-Control-Allow-Origin` header are rejected. +/// If `document_origin` is set, performs Same-Origin Policy and CORS checks. +/// For non-simple cross-origin requests, sends a preflight OPTIONS first. +/// Response headers are filtered per Access-Control-Expose-Headers. fn do_fetch( url_str: &str, method: &str, headers: &[(String, String)], body: Option<&[u8]>, document_origin: Option<&str>, + cors_mode: &str, + credentials_str: &str, ) -> Result { let url = we_url::Url::parse(url_str).map_err(|e| format!("Invalid URL: {e}"))?; @@ -216,51 +244,117 @@ fn do_fetch( other => return Err(format!("Unsupported HTTP method: {other}")), }; - let mut client = we_net::client::HttpClient::new(); - let response = client - .request(http_method, &url, &req_headers, body) - .map_err(|e| format!("Network error: {e}"))?; + let credentials_mode = parse_credentials_mode(credentials_str); - // Same-Origin Policy check: if we have a document origin, verify - // that cross-origin responses include a CORS header. - if let Some(doc_origin) = document_origin { + // Determine if this is a cross-origin request. + let is_cross_origin = if let Some(doc_origin) = document_origin { let resource_origin = url.origin(); let doc_parsed_origin = we_url::Url::parse(&format!("{doc_origin}/")) .map(|u| u.origin()) .unwrap_or(we_url::Origin::Opaque); + !doc_parsed_origin.same_origin(&resource_origin) + } else { + false + }; - if !doc_parsed_origin.same_origin(&resource_origin) { - // Cross-origin: check Access-Control-Allow-Origin. - let cors_allowed = response - .headers - .get("access-control-allow-origin") - .map(|v| { - let v = v.trim(); - v == "*" || v == doc_origin - }) - .unwrap_or(false); + let mut client = we_net::client::HttpClient::new(); - if !cors_allowed { - return Err(format!( - "Cross-origin request blocked: {url_str} (no CORS headers)" - )); - } + if is_cross_origin && cors_mode == "cors" { + let doc_origin = document_origin.unwrap(); + let doc_parsed_origin = we_url::Url::parse(&format!("{doc_origin}/")) + .map(|u| u.origin()) + .unwrap_or(we_url::Origin::Opaque); + + // Add Origin header. + req_headers.set("Origin", doc_origin); + + // Check if preflight is needed. + if we_net::cors::needs_preflight(http_method, &req_headers) { + let pf_headers = we_net::cors::build_preflight_headers( + &doc_parsed_origin, + http_method, + &req_headers, + ); + let pf_response = client + .request(we_net::http::Method::Options, &url, &pf_headers, None) + .map_err(|e| format!("Preflight error: {e}"))?; + + let pf_result = we_net::cors::validate_preflight( + &pf_response.headers, + &doc_parsed_origin, + credentials_mode, + ) + .map_err(|reason| format!("CORS preflight failed: {reason}"))?; + + we_net::cors::preflight_allows(&pf_result, http_method, &req_headers) + .map_err(|reason| format!("CORS preflight rejected: {reason}"))?; } - } - let resp_headers: Vec<(String, String)> = response - .headers - .iter() - .map(|(k, v)| (k.to_string(), v.to_string())) - .collect(); - - Ok(FetchResult { - status: response.status_code, - status_text: response.reason, - headers: resp_headers, - body: response.body, - url: url.serialize(), - }) + // Perform the actual request. + let response = client + .request(http_method, &url, &req_headers, body) + .map_err(|e| format!("Network error: {e}"))?; + + // Check CORS response headers. + let exposed = we_net::cors::check_cors_response( + &response.headers, + &doc_parsed_origin, + credentials_mode, + ) + .map_err(|reason| format!("Cross-origin request blocked: {reason}"))?; + + // Filter response headers to only exposed ones. + let filtered = we_net::cors::filter_response_headers(&response.headers, &exposed); + let resp_headers: Vec<(String, String)> = filtered + .iter() + .map(|(k, v)| (k.to_string(), v.to_string())) + .collect(); + + Ok(FetchResult { + status: response.status_code, + status_text: response.reason, + headers: resp_headers, + body: response.body, + url: url.serialize(), + }) + } else if is_cross_origin && cors_mode == "no-cors" { + // no-cors mode: make the request but return an opaque response. + let _response = client + .request(http_method, &url, &req_headers, body) + .map_err(|e| format!("Network error: {e}"))?; + + Ok(FetchResult { + status: 0, + status_text: String::new(), + headers: Vec::new(), + body: Vec::new(), + url: url.serialize(), + }) + } else if is_cross_origin { + // Default: block cross-origin if not in cors mode. + Err(format!( + "Cross-origin request blocked: {url_str} (mode is '{cors_mode}')" + )) + } else { + // Same-origin request. + let response = client + .request(http_method, &url, &req_headers, body) + .map_err(|e| format!("Network error: {e}"))?; + + let resp_headers: Vec<(String, String)> = response + .headers + .iter() + .map(|(k, v)| (k.to_string(), v.to_string())) + .collect(); + + Ok(FetchResult { + status: response.status_code, + status_text: response.reason, + headers: resp_headers, + body: response.body, + url: url.serialize(), + }) + } } // ── Response object creation ──────────────────────────────────── diff --git a/crates/net/src/cors.rs b/crates/net/src/cors.rs new file mode 100644 index 0000000..5ffa78c --- /dev/null +++ b/crates/net/src/cors.rs @@ -0,0 +1,1135 @@ +//! CORS (Cross-Origin Resource Sharing) per the Fetch Standard. +//! +//! Implements: +//! - Simple request detection (CORS-safelisted methods and headers) +//! - Preflight (OPTIONS) request construction and response validation +//! - Preflight result caching with max-age +//! - Access-Control-Allow-Origin / Allow-Credentials / Expose-Headers checks +//! - Credentials mode handling + +use std::collections::HashMap; +use std::time::{Duration, Instant}; + +use we_url::Origin; + +use crate::http::{Headers, Method}; + +// --------------------------------------------------------------------------- +// CORS request mode +// --------------------------------------------------------------------------- + +/// The CORS mode for a request. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CorsMode { + /// No CORS — request must be same-origin or the response is opaque. + NoCors, + /// CORS — cross-origin requests are allowed if the server opts in. + Cors, + /// Navigation — top-level document loads, always allowed. + Navigate, +} + +/// Credentials mode for a request. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CredentialsMode { + /// Never send credentials cross-origin. + Omit, + /// Send credentials only if same-origin. + SameOrigin, + /// Always send credentials (requires explicit server opt-in for cross-origin). + Include, +} + +// --------------------------------------------------------------------------- +// CORS-safelisted checks +// --------------------------------------------------------------------------- + +/// CORS-safelisted methods that do not trigger a preflight. +const SIMPLE_METHODS: &[Method] = &[Method::Get, Method::Head, Method::Post]; + +/// CORS-safelisted request header names (lowercase). +const SAFELISTED_HEADERS: &[&str] = &[ + "accept", + "accept-language", + "content-language", + "content-type", +]; + +/// Content-Type values that are CORS-safelisted. +const SAFELISTED_CONTENT_TYPES: &[&str] = &[ + "application/x-www-form-urlencoded", + "multipart/form-data", + "text/plain", +]; + +/// Check whether a method is CORS-safelisted (does not require preflight). +pub fn is_simple_method(method: Method) -> bool { + SIMPLE_METHODS.contains(&method) +} + +/// Check whether a header name is CORS-safelisted. +fn is_safelisted_header_name(name: &str) -> bool { + SAFELISTED_HEADERS.contains(&name.to_ascii_lowercase().as_str()) +} + +/// Check whether a Content-Type value is CORS-safelisted. +fn is_safelisted_content_type(value: &str) -> bool { + let mime = value + .split(';') + .next() + .unwrap_or("") + .trim() + .to_ascii_lowercase(); + SAFELISTED_CONTENT_TYPES.contains(&mime.as_str()) +} + +/// Determine whether a request requires a CORS preflight. +/// +/// A preflight is required if: +/// - The method is not GET, HEAD, or POST +/// - Any request header is not CORS-safelisted +/// - Content-Type (for POST) is not a safelisted value +pub fn needs_preflight(method: Method, headers: &Headers) -> bool { + if !is_simple_method(method) { + return true; + } + + for (name, value) in headers.iter() { + let lower = name.to_ascii_lowercase(); + if !is_safelisted_header_name(&lower) { + return true; + } + // Content-Type must also have a safelisted value. + if lower == "content-type" && !is_safelisted_content_type(value) { + return true; + } + } + + false +} + +// --------------------------------------------------------------------------- +// Preflight request construction +// --------------------------------------------------------------------------- + +/// Build the headers for a CORS preflight (OPTIONS) request. +pub fn build_preflight_headers( + origin: &Origin, + method: Method, + request_headers: &Headers, +) -> Headers { + let mut headers = Headers::new(); + headers.add("Origin", &origin.serialize()); + headers.add("Access-Control-Request-Method", method.as_str()); + + // Collect non-safelisted header names for Access-Control-Request-Headers. + let mut non_simple: Vec = Vec::new(); + for (name, value) in request_headers.iter() { + let lower = name.to_ascii_lowercase(); + if !is_safelisted_header_name(&lower) + || (lower == "content-type" && !is_safelisted_content_type(value)) + { + non_simple.push(lower); + } + } + if !non_simple.is_empty() { + non_simple.sort(); + non_simple.dedup(); + headers.add("Access-Control-Request-Headers", &non_simple.join(", ")); + } + + headers +} + +// --------------------------------------------------------------------------- +// Preflight response validation +// --------------------------------------------------------------------------- + +/// Result of validating a preflight response. +#[derive(Debug)] +pub struct PreflightResult { + /// Allowed methods from the preflight response. + pub allowed_methods: Vec, + /// Allowed headers from the preflight response. + pub allowed_headers: Vec, + /// Max-age for caching (seconds). Defaults to 5 seconds if not specified. + pub max_age: u64, +} + +/// Validate a preflight response and extract the allowed methods/headers. +/// +/// Returns `Err` with a human-readable reason if the preflight fails. +pub fn validate_preflight( + response_headers: &Headers, + request_origin: &Origin, + credentials_mode: CredentialsMode, +) -> Result { + // Check Access-Control-Allow-Origin. + let allow_origin = response_headers + .get("access-control-allow-origin") + .ok_or("preflight response missing Access-Control-Allow-Origin")?; + + let origin_str = request_origin.serialize(); + let allow_origin = allow_origin.trim(); + + if credentials_mode == CredentialsMode::Include { + // With credentials, wildcard is not allowed. + if allow_origin != origin_str { + return Err(format!( + "preflight: Access-Control-Allow-Origin must be '{origin_str}' \ + when credentials are included, got '{allow_origin}'" + )); + } + // Must also have Allow-Credentials: true. + let allow_creds = response_headers + .get("access-control-allow-credentials") + .unwrap_or(""); + if allow_creds.trim() != "true" { + return Err( + "preflight: Access-Control-Allow-Credentials must be 'true' \ + when credentials are included" + .to_string(), + ); + } + } else if allow_origin != "*" && allow_origin != origin_str { + return Err(format!( + "preflight: Access-Control-Allow-Origin '{allow_origin}' \ + does not match origin '{origin_str}'" + )); + } + + // Parse Access-Control-Allow-Methods. + let allowed_methods: Vec = response_headers + .get("access-control-allow-methods") + .unwrap_or("") + .split(',') + .map(|s| s.trim().to_uppercase()) + .filter(|s| !s.is_empty()) + .collect(); + + // Parse Access-Control-Allow-Headers. + let allowed_headers: Vec = response_headers + .get("access-control-allow-headers") + .unwrap_or("") + .split(',') + .map(|s| s.trim().to_ascii_lowercase()) + .filter(|s| !s.is_empty()) + .collect(); + + // Parse Access-Control-Max-Age (default to 5 seconds). + let max_age = response_headers + .get("access-control-max-age") + .and_then(|v| v.trim().parse::().ok()) + .unwrap_or(5); + + Ok(PreflightResult { + allowed_methods, + allowed_headers, + max_age, + }) +} + +/// Check whether the actual request's method and headers are allowed by +/// the preflight result. +pub fn preflight_allows( + result: &PreflightResult, + method: Method, + request_headers: &Headers, +) -> Result<(), String> { + let method_str = method.as_str().to_uppercase(); + + // Simple methods are always allowed even if not listed. + if !is_simple_method(method) + && !result.allowed_methods.iter().any(|m| m == &method_str) + && !result.allowed_methods.contains(&"*".to_string()) + { + return Err(format!( + "CORS: method {method_str} not allowed by preflight" + )); + } + + // Check non-safelisted headers. + for (name, value) in request_headers.iter() { + let lower = name.to_ascii_lowercase(); + if is_safelisted_header_name(&lower) { + if lower == "content-type" && !is_safelisted_content_type(value) { + // content-type with non-simple value needs to be allowed. + if !result.allowed_headers.contains(&lower) + && !result.allowed_headers.contains(&"*".to_string()) + { + return Err(format!( + "CORS: header '{lower}' with value '{value}' not allowed by preflight" + )); + } + } + continue; + } + if !result.allowed_headers.contains(&lower) + && !result.allowed_headers.contains(&"*".to_string()) + { + return Err(format!("CORS: header '{lower}' not allowed by preflight")); + } + } + + Ok(()) +} + +// --------------------------------------------------------------------------- +// CORS response check +// --------------------------------------------------------------------------- + +/// CORS-safelisted response header names that are always exposed. +const SAFELISTED_RESPONSE_HEADERS: &[&str] = &[ + "cache-control", + "content-language", + "content-length", + "content-type", + "expires", + "last-modified", + "pragma", +]; + +/// Check the Access-Control-Allow-Origin on the actual response. +/// +/// Returns `Ok(exposed_headers)` with the set of response header names +/// the script is allowed to read, or `Err` if the response is blocked. +pub fn check_cors_response( + response_headers: &Headers, + request_origin: &Origin, + credentials_mode: CredentialsMode, +) -> Result, String> { + let allow_origin = response_headers + .get("access-control-allow-origin") + .ok_or("CORS: response missing Access-Control-Allow-Origin")?; + + let origin_str = request_origin.serialize(); + let allow_origin = allow_origin.trim(); + + if credentials_mode == CredentialsMode::Include { + if allow_origin != origin_str { + return Err(format!( + "CORS: Access-Control-Allow-Origin must be '{origin_str}' \ + when credentials are included, got '{allow_origin}'" + )); + } + let allow_creds = response_headers + .get("access-control-allow-credentials") + .unwrap_or(""); + if allow_creds.trim() != "true" { + return Err("CORS: Access-Control-Allow-Credentials must be 'true' \ + when credentials are included" + .to_string()); + } + } else if allow_origin != "*" && allow_origin != origin_str { + return Err(format!( + "CORS: Access-Control-Allow-Origin '{allow_origin}' \ + does not match origin '{origin_str}'" + )); + } + + // Build list of exposed headers. + let mut exposed: Vec = SAFELISTED_RESPONSE_HEADERS + .iter() + .map(|s| s.to_string()) + .collect(); + + if let Some(expose_hdr) = response_headers.get("access-control-expose-headers") { + for name in expose_hdr.split(',') { + let name = name.trim().to_ascii_lowercase(); + if !name.is_empty() { + if name == "*" && credentials_mode != CredentialsMode::Include { + // Wildcard: expose all headers. + for (h, _) in response_headers.iter() { + let lower = h.to_ascii_lowercase(); + if !exposed.contains(&lower) { + exposed.push(lower); + } + } + break; + } + if !exposed.contains(&name) { + exposed.push(name); + } + } + } + } + + Ok(exposed) +} + +/// Filter response headers to only those the script is allowed to see. +pub fn filter_response_headers(headers: &Headers, exposed: &[String]) -> Headers { + let mut filtered = Headers::new(); + for (name, value) in headers.iter() { + let lower = name.to_ascii_lowercase(); + if exposed.contains(&lower) { + filtered.add(name, value); + } + } + filtered +} + +// --------------------------------------------------------------------------- +// Preflight cache +// --------------------------------------------------------------------------- + +/// Key for the preflight cache. +#[derive(Hash, Eq, PartialEq, Clone, Debug)] +struct PreflightCacheKey { + origin: String, + url: String, +} + +/// An entry in the preflight cache. +struct PreflightCacheEntry { + allowed_methods: Vec, + allowed_headers: Vec, + created: Instant, + max_age: Duration, +} + +/// Cache for CORS preflight results. +/// +/// Keyed by (request origin, URL). Entries expire after their max-age. +pub struct PreflightCache { + entries: HashMap, +} + +impl PreflightCache { + /// Create a new, empty preflight cache. + pub fn new() -> Self { + Self { + entries: HashMap::new(), + } + } + + /// Store a preflight result in the cache. + pub fn store(&mut self, origin: &Origin, url: &str, result: &PreflightResult) { + let key = PreflightCacheKey { + origin: origin.serialize(), + url: url.to_string(), + }; + self.entries.insert( + key, + PreflightCacheEntry { + allowed_methods: result.allowed_methods.clone(), + allowed_headers: result.allowed_headers.clone(), + created: Instant::now(), + max_age: Duration::from_secs(result.max_age), + }, + ); + } + + /// Look up a cached preflight result. + /// + /// Returns `Some` if a valid (non-expired) entry exists for the + /// given origin and URL, and the entry allows the requested method + /// and headers. + pub fn lookup( + &mut self, + origin: &Origin, + url: &str, + method: Method, + headers: &Headers, + ) -> bool { + let key = PreflightCacheKey { + origin: origin.serialize(), + url: url.to_string(), + }; + + let entry = match self.entries.get(&key) { + Some(e) => e, + None => return false, + }; + + // Check expiration. + if entry.created.elapsed() > entry.max_age { + self.entries.remove(&key); + return false; + } + + // Check method. + let method_str = method.as_str().to_uppercase(); + if !is_simple_method(method) + && !entry.allowed_methods.iter().any(|m| m == &method_str) + && !entry.allowed_methods.contains(&"*".to_string()) + { + return false; + } + + // Check non-safelisted headers. + for (name, value) in headers.iter() { + let lower = name.to_ascii_lowercase(); + if is_safelisted_header_name(&lower) { + if lower == "content-type" + && !is_safelisted_content_type(value) + && !entry.allowed_headers.contains(&lower) + && !entry.allowed_headers.contains(&"*".to_string()) + { + return false; + } + continue; + } + if !entry.allowed_headers.contains(&lower) + && !entry.allowed_headers.contains(&"*".to_string()) + { + return false; + } + } + + true + } + + /// Remove expired entries from the cache. + pub fn evict_expired(&mut self) { + self.entries.retain(|_, e| e.created.elapsed() <= e.max_age); + } +} + +impl Default for PreflightCache { + fn default() -> Self { + Self::new() + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +#[cfg(test)] +mod tests { + use super::*; + use we_url::Host; + + fn make_origin(scheme: &str, domain: &str, port: Option) -> Origin { + Origin::Tuple(scheme.to_string(), Host::Domain(domain.to_string()), port) + } + + // -- Simple method checks -- + + #[test] + fn get_is_simple() { + assert!(is_simple_method(Method::Get)); + } + + #[test] + fn head_is_simple() { + assert!(is_simple_method(Method::Head)); + } + + #[test] + fn post_is_simple() { + assert!(is_simple_method(Method::Post)); + } + + #[test] + fn put_is_not_simple() { + assert!(!is_simple_method(Method::Put)); + } + + #[test] + fn delete_is_not_simple() { + assert!(!is_simple_method(Method::Delete)); + } + + #[test] + fn patch_is_not_simple() { + assert!(!is_simple_method(Method::Patch)); + } + + #[test] + fn options_is_not_simple() { + assert!(!is_simple_method(Method::Options)); + } + + // -- needs_preflight -- + + #[test] + fn simple_get_no_preflight() { + let headers = Headers::new(); + assert!(!needs_preflight(Method::Get, &headers)); + } + + #[test] + fn simple_post_form_no_preflight() { + let mut headers = Headers::new(); + headers.add("Content-Type", "application/x-www-form-urlencoded"); + assert!(!needs_preflight(Method::Post, &headers)); + } + + #[test] + fn post_json_needs_preflight() { + let mut headers = Headers::new(); + headers.add("Content-Type", "application/json"); + assert!(needs_preflight(Method::Post, &headers)); + } + + #[test] + fn put_needs_preflight() { + let headers = Headers::new(); + assert!(needs_preflight(Method::Put, &headers)); + } + + #[test] + fn custom_header_needs_preflight() { + let mut headers = Headers::new(); + headers.add("X-Custom", "value"); + assert!(needs_preflight(Method::Get, &headers)); + } + + #[test] + fn accept_header_no_preflight() { + let mut headers = Headers::new(); + headers.add("Accept", "application/json"); + assert!(!needs_preflight(Method::Get, &headers)); + } + + #[test] + fn authorization_header_needs_preflight() { + let mut headers = Headers::new(); + headers.add("Authorization", "Bearer token123"); + assert!(needs_preflight(Method::Get, &headers)); + } + + // -- build_preflight_headers -- + + #[test] + fn preflight_headers_basic() { + let origin = make_origin("https", "example.com", None); + let mut req_headers = Headers::new(); + req_headers.add("X-Custom", "foo"); + + let pf = build_preflight_headers(&origin, Method::Put, &req_headers); + assert_eq!(pf.get("Origin").unwrap(), "https://example.com"); + assert_eq!(pf.get("Access-Control-Request-Method").unwrap(), "PUT"); + assert_eq!( + pf.get("Access-Control-Request-Headers").unwrap(), + "x-custom" + ); + } + + #[test] + fn preflight_headers_no_extra_headers() { + let origin = make_origin("https", "example.com", None); + let req_headers = Headers::new(); + + let pf = build_preflight_headers(&origin, Method::Delete, &req_headers); + assert_eq!(pf.get("Origin").unwrap(), "https://example.com"); + assert_eq!(pf.get("Access-Control-Request-Method").unwrap(), "DELETE"); + assert!(pf.get("Access-Control-Request-Headers").is_none()); + } + + #[test] + fn preflight_headers_multiple_custom() { + let origin = make_origin("https", "example.com", None); + let mut req_headers = Headers::new(); + req_headers.add("X-B", "2"); + req_headers.add("X-A", "1"); + + let pf = build_preflight_headers(&origin, Method::Post, &req_headers); + let requested = pf.get("Access-Control-Request-Headers").unwrap(); + // Should be sorted and deduplicated. + assert_eq!(requested, "x-a, x-b"); + } + + #[test] + fn preflight_headers_non_simple_content_type() { + let origin = make_origin("https", "example.com", None); + let mut req_headers = Headers::new(); + req_headers.add("Content-Type", "application/json"); + + let pf = build_preflight_headers(&origin, Method::Post, &req_headers); + assert_eq!( + pf.get("Access-Control-Request-Headers").unwrap(), + "content-type" + ); + } + + // -- validate_preflight -- + + #[test] + fn validate_preflight_wildcard_origin() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "*"); + headers.add("Access-Control-Allow-Methods", "PUT, DELETE"); + headers.add("Access-Control-Allow-Headers", "x-custom"); + headers.add("Access-Control-Max-Age", "600"); + + let result = validate_preflight(&headers, &origin, CredentialsMode::Omit).unwrap(); + assert!(result.allowed_methods.contains(&"PUT".to_string())); + assert!(result.allowed_methods.contains(&"DELETE".to_string())); + assert!(result.allowed_headers.contains(&"x-custom".to_string())); + assert_eq!(result.max_age, 600); + } + + #[test] + fn validate_preflight_exact_origin() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "https://example.com"); + headers.add("Access-Control-Allow-Methods", "POST"); + + let result = validate_preflight(&headers, &origin, CredentialsMode::Omit).unwrap(); + assert!(result.allowed_methods.contains(&"POST".to_string())); + } + + #[test] + fn validate_preflight_wrong_origin() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "https://other.com"); + + let result = validate_preflight(&headers, &origin, CredentialsMode::Omit); + assert!(result.is_err()); + } + + #[test] + fn validate_preflight_missing_allow_origin() { + let origin = make_origin("https", "example.com", None); + let headers = Headers::new(); + + let result = validate_preflight(&headers, &origin, CredentialsMode::Omit); + assert!(result.is_err()); + } + + #[test] + fn validate_preflight_credentials_wildcard_rejected() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "*"); + headers.add("Access-Control-Allow-Credentials", "true"); + + let result = validate_preflight(&headers, &origin, CredentialsMode::Include); + assert!(result.is_err()); + } + + #[test] + fn validate_preflight_credentials_exact_origin() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "https://example.com"); + headers.add("Access-Control-Allow-Credentials", "true"); + headers.add("Access-Control-Allow-Methods", "POST"); + + let result = validate_preflight(&headers, &origin, CredentialsMode::Include).unwrap(); + assert!(result.allowed_methods.contains(&"POST".to_string())); + } + + #[test] + fn validate_preflight_credentials_missing_allow_creds() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "https://example.com"); + + let result = validate_preflight(&headers, &origin, CredentialsMode::Include); + assert!(result.is_err()); + } + + #[test] + fn validate_preflight_default_max_age() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "*"); + + let result = validate_preflight(&headers, &origin, CredentialsMode::Omit).unwrap(); + assert_eq!(result.max_age, 5); + } + + // -- preflight_allows -- + + #[test] + fn preflight_allows_simple_method() { + let result = PreflightResult { + allowed_methods: vec![], + allowed_headers: vec![], + max_age: 5, + }; + // GET is always allowed even without being listed. + assert!(preflight_allows(&result, Method::Get, &Headers::new()).is_ok()); + } + + #[test] + fn preflight_allows_listed_method() { + let result = PreflightResult { + allowed_methods: vec!["PUT".to_string()], + allowed_headers: vec![], + max_age: 5, + }; + assert!(preflight_allows(&result, Method::Put, &Headers::new()).is_ok()); + } + + #[test] + fn preflight_rejects_unlisted_method() { + let result = PreflightResult { + allowed_methods: vec!["PUT".to_string()], + allowed_headers: vec![], + max_age: 5, + }; + assert!(preflight_allows(&result, Method::Delete, &Headers::new()).is_err()); + } + + #[test] + fn preflight_allows_listed_header() { + let result = PreflightResult { + allowed_methods: vec![], + allowed_headers: vec!["x-custom".to_string()], + max_age: 5, + }; + let mut headers = Headers::new(); + headers.add("X-Custom", "value"); + assert!(preflight_allows(&result, Method::Get, &headers).is_ok()); + } + + #[test] + fn preflight_rejects_unlisted_header() { + let result = PreflightResult { + allowed_methods: vec![], + allowed_headers: vec!["x-other".to_string()], + max_age: 5, + }; + let mut headers = Headers::new(); + headers.add("X-Custom", "value"); + assert!(preflight_allows(&result, Method::Get, &headers).is_err()); + } + + #[test] + fn preflight_allows_wildcard_method() { + let result = PreflightResult { + allowed_methods: vec!["*".to_string()], + allowed_headers: vec![], + max_age: 5, + }; + assert!(preflight_allows(&result, Method::Delete, &Headers::new()).is_ok()); + } + + #[test] + fn preflight_allows_wildcard_header() { + let result = PreflightResult { + allowed_methods: vec![], + allowed_headers: vec!["*".to_string()], + max_age: 5, + }; + let mut headers = Headers::new(); + headers.add("X-Anything", "yes"); + assert!(preflight_allows(&result, Method::Get, &headers).is_ok()); + } + + // -- check_cors_response -- + + #[test] + fn cors_response_wildcard_origin() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "*"); + headers.add("Content-Type", "text/plain"); + headers.add("X-Custom", "hidden"); + + let exposed = check_cors_response(&headers, &origin, CredentialsMode::Omit).unwrap(); + assert!(exposed.contains(&"content-type".to_string())); + // X-Custom is not exposed without Expose-Headers. + assert!(!exposed.contains(&"x-custom".to_string())); + } + + #[test] + fn cors_response_exact_origin() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "https://example.com"); + + let exposed = check_cors_response(&headers, &origin, CredentialsMode::Omit).unwrap(); + assert!(exposed.contains(&"content-type".to_string())); + } + + #[test] + fn cors_response_wrong_origin_blocked() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "https://other.com"); + + let result = check_cors_response(&headers, &origin, CredentialsMode::Omit); + assert!(result.is_err()); + } + + #[test] + fn cors_response_no_header_blocked() { + let origin = make_origin("https", "example.com", None); + let headers = Headers::new(); + + let result = check_cors_response(&headers, &origin, CredentialsMode::Omit); + assert!(result.is_err()); + } + + #[test] + fn cors_response_expose_headers() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "*"); + headers.add("Access-Control-Expose-Headers", "X-Custom, X-Request-Id"); + headers.add("X-Custom", "visible"); + headers.add("X-Request-Id", "abc"); + headers.add("X-Secret", "hidden"); + + let exposed = check_cors_response(&headers, &origin, CredentialsMode::Omit).unwrap(); + assert!(exposed.contains(&"x-custom".to_string())); + assert!(exposed.contains(&"x-request-id".to_string())); + assert!(!exposed.contains(&"x-secret".to_string())); + } + + #[test] + fn cors_response_expose_wildcard() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "*"); + headers.add("Access-Control-Expose-Headers", "*"); + headers.add("X-Custom", "visible"); + + let exposed = check_cors_response(&headers, &origin, CredentialsMode::Omit).unwrap(); + assert!(exposed.contains(&"x-custom".to_string())); + } + + #[test] + fn cors_response_credentials_wildcard_rejected() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "*"); + headers.add("Access-Control-Allow-Credentials", "true"); + + let result = check_cors_response(&headers, &origin, CredentialsMode::Include); + assert!(result.is_err()); + } + + #[test] + fn cors_response_credentials_ok() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "https://example.com"); + headers.add("Access-Control-Allow-Credentials", "true"); + + let exposed = check_cors_response(&headers, &origin, CredentialsMode::Include).unwrap(); + assert!(exposed.contains(&"content-type".to_string())); + } + + // -- filter_response_headers -- + + #[test] + fn filter_headers_basic() { + let mut headers = Headers::new(); + headers.add("Content-Type", "text/plain"); + headers.add("X-Custom", "value"); + headers.add("X-Secret", "hidden"); + + let exposed = vec!["content-type".to_string(), "x-custom".to_string()]; + let filtered = filter_response_headers(&headers, &exposed); + assert!(filtered.get("Content-Type").is_some()); + assert!(filtered.get("X-Custom").is_some()); + assert!(filtered.get("X-Secret").is_none()); + } + + // -- PreflightCache -- + + #[test] + fn cache_store_and_lookup() { + let origin = make_origin("https", "example.com", None); + let mut cache = PreflightCache::new(); + + let result = PreflightResult { + allowed_methods: vec!["PUT".to_string()], + allowed_headers: vec!["x-custom".to_string()], + max_age: 600, + }; + cache.store(&origin, "https://api.example.com/data", &result); + + let mut headers = Headers::new(); + headers.add("X-Custom", "value"); + + assert!(cache.lookup( + &origin, + "https://api.example.com/data", + Method::Put, + &headers, + )); + } + + #[test] + fn cache_miss_different_url() { + let origin = make_origin("https", "example.com", None); + let mut cache = PreflightCache::new(); + + let result = PreflightResult { + allowed_methods: vec!["PUT".to_string()], + allowed_headers: vec![], + max_age: 600, + }; + cache.store(&origin, "https://api.example.com/a", &result); + + assert!(!cache.lookup( + &origin, + "https://api.example.com/b", + Method::Put, + &Headers::new(), + )); + } + + #[test] + fn cache_miss_different_origin() { + let origin_a = make_origin("https", "a.com", None); + let origin_b = make_origin("https", "b.com", None); + let mut cache = PreflightCache::new(); + + let result = PreflightResult { + allowed_methods: vec!["PUT".to_string()], + allowed_headers: vec![], + max_age: 600, + }; + cache.store(&origin_a, "https://api.example.com/data", &result); + + assert!(!cache.lookup( + &origin_b, + "https://api.example.com/data", + Method::Put, + &Headers::new(), + )); + } + + #[test] + fn cache_miss_disallowed_method() { + let origin = make_origin("https", "example.com", None); + let mut cache = PreflightCache::new(); + + let result = PreflightResult { + allowed_methods: vec!["PUT".to_string()], + allowed_headers: vec![], + max_age: 600, + }; + cache.store(&origin, "https://api.example.com/data", &result); + + assert!(!cache.lookup( + &origin, + "https://api.example.com/data", + Method::Delete, + &Headers::new(), + )); + } + + #[test] + fn cache_miss_disallowed_header() { + let origin = make_origin("https", "example.com", None); + let mut cache = PreflightCache::new(); + + let result = PreflightResult { + allowed_methods: vec!["PUT".to_string()], + allowed_headers: vec!["x-allowed".to_string()], + max_age: 600, + }; + cache.store(&origin, "https://api.example.com/data", &result); + + let mut headers = Headers::new(); + headers.add("X-Disallowed", "value"); + + assert!(!cache.lookup( + &origin, + "https://api.example.com/data", + Method::Put, + &headers, + )); + } + + #[test] + fn cache_allows_simple_method_without_listing() { + let origin = make_origin("https", "example.com", None); + let mut cache = PreflightCache::new(); + + let result = PreflightResult { + allowed_methods: vec![], + allowed_headers: vec!["x-custom".to_string()], + max_age: 600, + }; + cache.store(&origin, "https://api.example.com/data", &result); + + let mut headers = Headers::new(); + headers.add("X-Custom", "value"); + + // GET is a simple method, should be allowed even without listing. + assert!(cache.lookup( + &origin, + "https://api.example.com/data", + Method::Get, + &headers, + )); + } + + #[test] + fn cache_evict_expired() { + let origin = make_origin("https", "example.com", None); + let mut cache = PreflightCache::new(); + + // Store with max_age=0 so it's immediately expired. + let result = PreflightResult { + allowed_methods: vec!["PUT".to_string()], + allowed_headers: vec![], + max_age: 0, + }; + cache.store(&origin, "https://api.example.com/data", &result); + + // Wait a tiny bit to ensure expiration. + std::thread::sleep(std::time::Duration::from_millis(10)); + + cache.evict_expired(); + assert!(cache.entries.is_empty()); + } + + #[test] + fn default_cache_is_empty() { + let cache = PreflightCache::default(); + assert!(cache.entries.is_empty()); + } + + // -- CORS-safelisted content type checks -- + + #[test] + fn form_urlencoded_is_safelisted() { + assert!(is_safelisted_content_type( + "application/x-www-form-urlencoded" + )); + } + + #[test] + fn multipart_is_safelisted() { + assert!(is_safelisted_content_type("multipart/form-data")); + } + + #[test] + fn text_plain_is_safelisted() { + assert!(is_safelisted_content_type("text/plain")); + } + + #[test] + fn json_is_not_safelisted() { + assert!(!is_safelisted_content_type("application/json")); + } + + #[test] + fn content_type_with_params_is_safelisted() { + assert!(is_safelisted_content_type("text/plain; charset=utf-8")); + } + + // -- Safelisted response headers -- + + #[test] + fn safelisted_response_headers_always_exposed() { + let origin = make_origin("https", "example.com", None); + let mut headers = Headers::new(); + headers.add("Access-Control-Allow-Origin", "*"); + headers.add("Content-Type", "text/html"); + headers.add("Cache-Control", "max-age=3600"); + + let exposed = check_cors_response(&headers, &origin, CredentialsMode::Omit).unwrap(); + assert!(exposed.contains(&"content-type".to_string())); + assert!(exposed.contains(&"cache-control".to_string())); + assert!(exposed.contains(&"content-length".to_string())); + assert!(exposed.contains(&"expires".to_string())); + assert!(exposed.contains(&"last-modified".to_string())); + assert!(exposed.contains(&"pragma".to_string())); + } +} diff --git a/crates/net/src/lib.rs b/crates/net/src/lib.rs index 04c600c..1684fef 100644 --- a/crates/net/src/lib.rs +++ b/crates/net/src/lib.rs @@ -1,6 +1,7 @@ -//! TCP, DNS, pure-Rust TLS 1.3, HTTP/1.1, HTTP/2. +//! TCP, DNS, pure-Rust TLS 1.3, HTTP/1.1, HTTP/2, CORS. pub mod client; +pub mod cors; pub mod dns; pub mod http; pub mod tcp;