use super::*;
#[cfg(feature = "handlers")]
#[allow(clippy::too_many_arguments)]
pub(super) async fn dispatch_handler(
runtime: &HandlerRuntime,
deploy: &DeployStore,
manifest: &Manifest,
project: &str,
site: &str,
request_path: &str,
site_config: Option<&SiteConfig>,
handler: &boatramp_core::config::HandlerConfig,
mut request: Request,
client_ip: IpAddr,
preview: Option<&str>,
) -> Response {
let Some(inner) = runtime.inner.as_ref() else {
return not_found();
};
let project_ref = boatramp_core::project::ProjectRef::new(project);
let base = match preview {
Some(id) => format!("{site}/_preview/{id}"),
None => site.to_string(),
};
let scope = project_ref.qualified(&base);
set_forwarded_headers(&mut request, client_ip);
rewrite_request_uri(&mut request, request_path);
let Some(site_handlers) = site_config
.and_then(|c| c.handlers.as_ref())
.filter(|h| h.enabled)
else {
return not_found();
};
match cookie_auth_outcome(request.headers(), site_handlers.cookie_auth.as_ref()) {
CookieAuthOutcome::None => {}
CookieAuthOutcome::Reject => {
return (
StatusCode::FORBIDDEN,
"cross-origin cookie-authenticated request rejected\n",
)
.into_response();
}
CookieAuthOutcome::Inject(token) => {
if let Ok(value) = HeaderValue::try_from(format!("Bearer {token}")) {
request.headers_mut().insert(header::AUTHORIZATION, value);
}
}
}
if let Some(gql) = site_handlers.graphql.as_ref().filter(|g| g.enabled) {
if gql.graphiql && request.method() == Method::GET {
let wants_html = request
.headers()
.get(header::ACCEPT)
.and_then(|v| v.to_str().ok())
.is_some_and(|a| a.contains("text/html"));
if wants_html {
return graphql_graphiql::page();
}
}
if request.method() == Method::POST {
let content_type = request
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.map(str::to_string);
let is_upload = content_type.as_deref().is_some_and(|ct| {
ct.contains("multipart/form-data")
|| ct.contains("application/x-www-form-urlencoded")
});
if !is_upload {
let (parts, body) = request.into_parts();
let mut body_bytes =
match axum::body::to_bytes(body, graphql_guard::MAX_QUERY_BYTES).await {
Ok(raw) => raw.to_vec(),
Err(_) => return graphql_guard::too_large_response(),
};
let mut effective_query: Option<String> = None;
let mut variables = serde_json::Value::Object(Default::default());
if let Ok(mut json) = serde_json::from_slice::<serde_json::Value>(&body_bytes) {
if let Some(vars) = json.get("variables").filter(|v| v.is_object()) {
variables = vars.clone();
}
if gql.persisted_queries || gql.safelist {
match graphql_apq::resolve_stored(
inner.kv.as_ref(),
&scope,
&json,
gql.safelist,
)
.await
{
graphql_apq::Resolved::Error(msg) => {
return graphql_apq::error_response(&msg)
}
graphql_apq::Resolved::Query(q) => {
json["query"] = serde_json::Value::String(q.clone());
if let Ok(v) = serde_json::to_vec(&json) {
body_bytes = v;
}
effective_query = Some(q);
}
graphql_apq::Resolved::Passthrough => {}
}
}
if effective_query.is_none() {
effective_query = json
.get("query")
.and_then(|v| v.as_str())
.map(str::to_string);
}
} else {
effective_query =
graphql_guard::query_from_body(content_type.as_deref(), &body_bytes);
}
if let Some(query) = &effective_query {
if let graphql_guard::GuardVerdict::Reject(reason) =
graphql_guard::guard_query(query, &graphql_guard::limits_from(gql))
{
return graphql_guard::error_response(&reason);
}
let bearer = parts
.headers
.get(header::AUTHORIZATION)
.and_then(|v| v.to_str().ok())
.and_then(|s| {
s.strip_prefix("Bearer ")
.or_else(|| s.strip_prefix("bearer "))
})
.map(str::to_string);
if let Some(topic) = graphql_subscription::subscription_topic(query) {
let after = parts
.headers
.get("last-event-id")
.and_then(|v| v.to_str().ok())
.map(str::to_string);
return crate::stream::serve_graphql_subscription(
inner,
site,
site_handlers,
&topic,
after,
client_ip,
preview,
)
.await;
}
if gql.federated {
return federation_gateway(
inner,
project,
query,
&variables,
bearer.as_deref(),
)
.await;
}
if let Some(data) = gql.data.as_ref().filter(|d| d.enabled) {
return data_connector_serve(
inner,
project,
site,
data,
query,
&variables,
bearer.as_deref(),
)
.await;
}
}
request = Request::from_parts(parts, axum::body::Body::from(body_bytes));
}
}
}
let cache_cfg = handler_cache::config_for(site_handlers);
let cache_key = cache_cfg.as_ref().and_then(|cfg| {
handler_cache::request_lookupable(cfg, request.method()).then(|| {
let path_and_query = request.uri().path_and_query().map_or("/", |pq| pq.as_str());
handler_cache::cache_key(&scope, request.method(), path_and_query)
})
});
if let Some(key) = &cache_key {
if let Some(hit) = handler_cache::lookup_response(
inner.kv.as_ref(),
key,
request.headers(),
handler_cache::now_secs(),
)
.await
{
return hit;
}
}
let cache_write = match (&cache_cfg, &cache_key) {
(Some(cfg), Some(key)) => Some((
cfg.clone(),
key.clone(),
request.method().clone(),
request.headers().clone(),
)),
_ => None,
};
let Some(entry) = manifest.files.get(&handler.component) else {
tracing::warn!(site, component = %handler.component, "handler component missing from deployment");
return handler_unavailable();
};
let wasm = match read_blob_fully(deploy, &entry.hash).await {
Ok(bytes) => bytes,
Err(response) => return response,
};
let request_id = request
.extensions()
.get::<crate::RequestId>()
.map(|r| r.0.clone());
let bindings = build_bindings(
inner,
boatramp_core::project::ProjectRef::new(project),
site,
&scope,
preview,
&handler.imports,
site_handlers,
&handler.env,
&handler.invoke_targets,
0,
request_id.as_deref(),
)
.await;
let _site_permit = match acquire_site_permit(inner, &scope, site_handlers) {
Ok(permit) => permit,
Err(()) => {
return (
StatusCode::SERVICE_UNAVAILABLE,
"site handler concurrency limit reached\n",
)
.into_response()
}
};
let limits = effective_limits(site_handlers, handler);
let start = std::time::Instant::now();
let result = inner
.engine
.serve_with_limits(&entry.hash, &wasm, request, bindings, limits)
.await;
inner.metrics.observe(
site,
metrics::Trigger::Http,
&handler.route,
&entry.hash,
metrics::Outcome::from_result(&result),
start.elapsed(),
);
match result {
Ok(response) => {
let (parts, body) = response.into_parts();
let response = axum::http::Response::from_parts(parts, axum::body::Body::new(body));
match &cache_write {
Some((cfg, key, method, req_headers)) => {
handler_cache::maybe_store(
inner.kv.clone(),
cfg,
key,
method,
req_headers,
response,
handler_cache::now_secs(),
)
.await
}
None => response,
}
}
Err(err) => {
tracing::warn!(site, route = %handler.route, %err, "handler invocation failed");
handler_error_response(&err)
}
}
}
#[cfg(feature = "handlers")]
async fn federation_gateway(
inner: &HandlerRuntimeInner,
project: &str,
query: &str,
variables: &serde_json::Value,
bearer: Option<&str>,
) -> Response {
let cached = match inner
.graphql_cache
.supergraph(inner.kv.as_ref(), project)
.await
{
Ok(c) => c,
Err(err) => {
return (
StatusCode::BAD_GATEWAY,
format!("supergraph composition failed: {err}\n"),
)
.into_response()
}
};
let op_hash = crate::graphql_apq::sha256_hex(query);
let plan =
match inner
.graphql_cache
.plan(project, cached.version, &op_hash, query, &cached.supergraph)
{
Ok(plan) => plan,
Err(_) => {
return graphql_guard::error_response(
"the query cannot be planned against the supergraph",
)
}
};
let Some(invoker) = inner.invoker.get() else {
return (
StatusCode::SERVICE_UNAVAILABLE,
"federation gateway: no invoker configured\n",
)
.into_response();
};
let sql_subgraphs = (*cached.sql_subgraphs).clone();
let runner = crate::graphql_gateway::BackendRouter::new(
invoker.scoped(boatramp_core::project::ProjectRef::new(project)),
project.to_string(),
inner.sql.clone(),
sql_subgraphs,
bearer.map(str::to_string),
);
axum::Json(crate::graphql_gateway::execute(&plan, &runner, variables).await).into_response()
}
#[cfg(feature = "handlers")]
async fn data_connector_serve(
inner: &HandlerRuntimeInner,
project: &str,
site: &str,
cfg: &boatramp_core::config::HandlerGraphqlDataConfig,
query: &str,
variables: &serde_json::Value,
bearer: Option<&str>,
) -> Response {
let Some(provider) = &inner.sql else {
return (
StatusCode::SERVICE_UNAVAILABLE,
"graphql data connector: this server has no SQL backend configured\n",
)
.into_response();
};
let backend = match provider.database(project, site, &cfg.source).await {
Ok(backend) => backend,
Err(err) => {
tracing::warn!(site, %err, "graphql data connector: opening the database failed");
return (
StatusCode::BAD_GATEWAY,
"graphql data connector: database unavailable\n",
)
.into_response();
}
};
let schema = match crate::graphql_data::introspect::introspect_sqlite(backend.as_ref()).await {
Ok(schema) => schema,
Err(err) => {
tracing::warn!(site, %err, "graphql data connector: introspection failed");
return (
StatusCode::BAD_GATEWAY,
"graphql data connector: schema introspection failed\n",
)
.into_response();
}
};
let policy = crate::graphql_data::policy_from_config(cfg);
let claims = crate::graphql_data::request_claims(project, bearer, cfg).await;
let dialect = crate::graphql_data::dialect::Sqlite;
let response = if crate::graphql_data::compile::is_mutation(query) {
if !cfg.mutations {
serde_json::json!({ "errors": [ { "message": "mutations are not enabled for this endpoint" } ] })
} else {
crate::graphql_data::runner::execute_mutation(
backend.as_ref(),
&dialect,
&schema,
&policy,
&claims,
query,
variables,
)
.await
}
} else {
let invoker = inner
.invoker
.get()
.map(|inv| inv.scoped(boatramp_core::project::ProjectRef::new(project)));
crate::graphql_data::runner::execute(
backend.as_ref(),
&dialect,
&schema,
&policy,
&claims,
query,
variables,
invoker.as_deref(),
bearer,
0, )
.await
};
axum::Json(response).into_response()
}
#[cfg(feature = "handlers")]
pub(super) fn set_forwarded_headers(request: &mut Request, client_ip: IpAddr) {
let headers = request.headers_mut();
if let Ok(value) = HeaderValue::from_str(&client_ip.to_string()) {
headers.insert(HeaderName::from_static("x-forwarded-for"), value);
}
if let Some(host) = headers.get(header::HOST).cloned() {
headers.insert(HeaderName::from_static("x-forwarded-host"), host);
}
if !headers.contains_key("x-forwarded-proto") {
headers.insert(
HeaderName::from_static("x-forwarded-proto"),
HeaderValue::from_static("http"),
);
}
}
#[cfg(feature = "handlers")]
fn rewrite_request_uri(request: &mut Request, request_path: &str) {
let authority = request
.headers()
.get(header::HOST)
.and_then(|value| value.to_str().ok())
.filter(|host| !host.is_empty())
.unwrap_or("localhost")
.to_string();
let path_and_query = match request.uri().query() {
Some(query) => format!("{request_path}?{query}"),
None => request_path.to_string(),
};
if let Ok(uri) = format!("http://{authority}{path_and_query}").parse() {
*request.uri_mut() = uri;
}
}
#[cfg(feature = "handlers")]
#[allow(clippy::too_many_arguments)]
pub(super) async fn precheck_component(
deploy: &DeployStore,
manifest: &Manifest,
site_handlers: &boatramp_core::config::HandlersSiteConfig,
inner: &HandlerRuntimeInner,
max_component: u64,
imports: &[String],
component: &str,
label: &str,
is_consumer: bool,
) -> Result<(), String> {
for import in imports {
if !site_handlers.allow_imports.iter().any(|a| a == import) {
return Err(format!(
"{label} requests import {import:?} the site does not allow"
));
}
if import == "sql" && inner.sql.is_none() {
return Err(format!(
"{label} requests `sql` but this server has no SQL backend configured"
));
}
if import == "wasi:messaging" && inner.messaging.is_none() {
return Err(format!(
"{label} requests `wasi:messaging` but this server has no messaging backend"
));
}
}
let entry = manifest
.files
.get(component)
.ok_or_else(|| format!("{label} component {component:?} missing from deployment"))?;
if max_component != 0 && entry.size > max_component {
return Err(format!(
"{label} component {component:?} is {} bytes, over the {max_component}-byte limit",
entry.size
));
}
let wasm = read_blob_bytes(deploy, &entry.hash)
.await
.map_err(|err| format!("reading {label} component: {err}"))?;
if is_consumer {
inner
.engine
.precompile_consumer(&entry.hash, &wasm)
.map_err(|err| format!("{label} is not a valid wasi:messaging consumer: {err}"))?;
} else {
inner
.engine
.precompile(&entry.hash, &wasm)
.map_err(|err| format!("{label} failed to compile: {err}"))?;
}
Ok(())
}
#[cfg(feature = "handlers")]
pub(super) async fn read_blob_bytes(
deploy: &DeployStore,
hash: &str,
) -> Result<Vec<u8>, DeployError> {
let object = deploy.open_blob(hash).await?;
let mut body = object.body;
let mut buf = Vec::new();
while let Some(chunk) = body.next().await {
buf.extend_from_slice(&chunk?);
}
Ok(buf)
}
#[cfg(feature = "handlers")]
pub(super) async fn read_blob_fully(deploy: &DeployStore, hash: &str) -> Result<Vec<u8>, Response> {
read_blob_bytes(deploy, hash)
.await
.map_err(deploy_error_response)
}
enum CookieAuthOutcome {
None,
Inject(String),
Reject,
}
fn cookie_auth_outcome(
headers: &HeaderMap,
cookie_auth: Option<&boatramp_core::config::CookieAuthConfig>,
) -> CookieAuthOutcome {
let Some(cookie_auth) = cookie_auth else {
return CookieAuthOutcome::None;
};
if headers.contains_key(header::AUTHORIZATION) {
return CookieAuthOutcome::None;
}
let Some(token) = cookie_value(headers, &cookie_auth.cookie_name) else {
return CookieAuthOutcome::None;
};
if !origin_allowed(headers, &cookie_auth.allowed_origins) {
return CookieAuthOutcome::Reject;
}
CookieAuthOutcome::Inject(token)
}
fn cookie_value(headers: &HeaderMap, name: &str) -> Option<String> {
let cookies = headers.get(header::COOKIE)?.to_str().ok()?;
cookies.split(';').find_map(|pair| {
let (k, v) = pair.trim().split_once('=')?;
(k == name).then(|| v.trim().to_string())
})
}
fn referer_origin(referer: &str) -> Option<String> {
let (scheme, rest) = referer.split_once("://")?;
let authority = rest
.split(['/', '?', '#'])
.next()
.filter(|a| !a.is_empty())?;
Some(format!("{scheme}://{authority}"))
}
fn origin_allowed(headers: &HeaderMap, allowed: &[String]) -> bool {
let origin = headers
.get(header::ORIGIN)
.and_then(|v| v.to_str().ok())
.map(str::to_string)
.or_else(|| {
headers
.get(header::REFERER)
.and_then(|v| v.to_str().ok())
.and_then(referer_origin)
});
match origin {
None => true,
Some(origin) => is_same_origin(headers, &origin) || allowed.iter().any(|a| a == &origin),
}
}
fn is_same_origin(headers: &HeaderMap, origin: &str) -> bool {
let Some(host) = headers.get(header::HOST).and_then(|v| v.to_str().ok()) else {
return false;
};
let origin_authority = origin.split_once("://").map_or(origin, |(_, a)| a);
!host.is_empty() && origin_authority.eq_ignore_ascii_case(host)
}
#[cfg(feature = "handlers")]
pub(super) fn granted_sql_databases(imports: &[String], allow_imports: &[String]) -> Vec<String> {
let has = |list: &[String], v: &str| list.iter().any(|i| i == v);
let mut names: Vec<String> = Vec::new();
if has(imports, "sql") && has(allow_imports, "sql") {
names.push(String::new()); }
let handler_wildcard = has(imports, "sql:*");
for allowed in allow_imports {
let Some(name) = allowed.strip_prefix("sql:") else {
continue;
};
if name.is_empty() || name == "*" {
continue; }
if handler_wildcard || has(imports, allowed) {
names.push(name.to_string());
}
}
names
}
#[cfg(feature = "handlers")]
#[allow(clippy::too_many_arguments)]
pub(super) async fn build_bindings(
inner: &HandlerRuntimeInner,
project: boatramp_core::project::ProjectRef<'_>,
site: &str,
scope: &str,
preview: Option<&str>,
imports: &[String],
site_handlers: &boatramp_core::config::HandlersSiteConfig,
deploy_env: &std::collections::BTreeMap<String, String>,
invoke_targets: &[String],
depth: u32,
request_id: Option<&str>,
) -> boatramp_handlers::Bindings {
let granted = |name: &str| {
imports.iter().any(|i| i == name) && site_handlers.allow_imports.iter().any(|a| a == name)
};
let mut bindings = boatramp_handlers::Bindings::new(scope);
if granted("wasi:keyvalue") {
bindings = bindings.with_keyvalue(scope, inner.kv.clone());
}
if granted("wasi:blobstore") {
let max_blob = inner.max_blob_bytes.get().copied().unwrap_or(0);
bindings = bindings.with_blobstore(scope, inner.storage.clone(), max_blob);
}
if let Some(provider) = &inner.sql {
for name in granted_sql_databases(imports, &site_handlers.allow_imports) {
let opened = match preview {
Some(id) => {
provider
.preview_database(project.as_str(), site, &name, id)
.await
}
None => provider.database(project.as_str(), site, &name).await,
};
match opened {
Ok(backend) => bindings = bindings.with_sql(name.clone(), backend),
Err(err) => {
tracing::warn!(site, database = %name, %err, "opening SQL database failed");
}
}
}
}
if granted("wasi:messaging") {
if let Some(messaging) = &inner.messaging {
bindings = bindings.with_messaging(
format!("{scope}/"),
format!("{}/", project.qualified("bus")),
messaging.clone(),
);
}
}
if granted("invoke") && !invoke_targets.is_empty() {
if let Some(invoker) = inner.invoker.get() {
bindings =
bindings.with_invoke(invoker.scoped(project), invoke_targets.to_vec(), depth);
}
}
if granted("graphql") {
if let Some(runner) = inner.federation_runner.get() {
bindings = bindings.with_graphql(runner.scoped(project), depth);
}
}
if !site_handlers.disable_log_capture {
inner.logs.configure(site, site_handlers.max_log_rate);
bindings = bindings.with_logging(
site.to_string(),
request_id.map(str::to_string),
inner.logs.clone(),
);
}
bindings = bindings.with_env(resolve_env(site, deploy_env, site_handlers));
bindings
}
#[cfg(feature = "handlers")]
pub(super) fn resolve_env(
site: &str,
deploy_env: &std::collections::BTreeMap<String, String>,
site_handlers: &boatramp_core::config::HandlersSiteConfig,
) -> Vec<(String, String)> {
let mut env: Vec<(String, String)> = deploy_env
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
for (guest_name, host_ref) in &site_handlers.secrets {
match std::env::var(host_ref) {
Ok(value) => {
env.retain(|(k, _)| k != guest_name);
env.push((guest_name.clone(), value));
}
Err(_) => tracing::warn!(
site,
secret = %guest_name,
"site secret references env var {host_ref}, which is not set; not injected"
),
}
}
env
}
#[cfg(feature = "handlers")]
#[allow(clippy::too_many_arguments)]
pub(super) async fn dispatch_consumer_batch(
engine: &boatramp_handlers::HandlerEngine,
messaging: &dyn boatramp_core::messaging::Messaging,
metrics: &metrics::Metrics,
site: &str,
namespaced_topic: &str,
scope_prefix: &str,
group: &str,
start: boatramp_core::messaging::StartPosition,
component_hash: &str,
component: &[u8],
bindings: &boatramp_handlers::Bindings,
limits: boatramp_handlers::Limits,
lease: Duration,
max_attempts: u32,
batch: usize,
) -> usize {
let claimed = match messaging
.claim_grouped(namespaced_topic, group, start, lease, batch, max_attempts)
.await
{
Ok(claimed) => claimed,
Err(err) => {
tracing::warn!(topic = namespaced_topic, %err, "messaging claim failed");
return 0;
}
};
let mut acked = 0;
for msg in claimed {
let guest_topic = msg.topic.strip_prefix(scope_prefix).unwrap_or(&msg.topic);
let start = std::time::Instant::now();
let result = engine
.dispatch_message(
component_hash,
component,
guest_topic,
&msg.payload,
bindings.clone(),
limits,
)
.await;
metrics.observe(
site,
metrics::Trigger::Consumer,
guest_topic,
component_hash,
metrics::Outcome::from_result(&result),
start.elapsed(),
);
match result {
Ok(()) => match messaging.ack(&msg).await {
Ok(()) => acked += 1,
Err(err) => tracing::warn!(id = msg.id, %err, "messaging ack failed"),
},
Err(err) => {
tracing::warn!(
id = msg.id,
attempts = msg.attempts,
%err,
"consumer failed; redelivering (dead-letters after max attempts)"
);
let _ = messaging.nack(&msg).await;
}
}
}
acked
}
#[cfg(all(test, feature = "handlers"))]
mod vhost_tests {
use super::*;
#[test]
fn a_wildcard_routed_request_carries_the_real_public_host_to_the_guest() {
let mut req = Request::builder()
.method("GET")
.uri("/_sites/portal/dashboard?tenant=7")
.header("host", "tenant7.construens.com")
.body(Body::empty())
.unwrap();
set_forwarded_headers(&mut req, std::net::IpAddr::from([203, 0, 113, 9]));
rewrite_request_uri(&mut req, "/dashboard");
assert_eq!(req.uri().host(), Some("tenant7.construens.com"));
assert_eq!(req.uri().path(), "/dashboard");
assert_eq!(req.uri().query(), Some("tenant=7"));
assert_eq!(
req.headers().get(header::HOST).unwrap(),
"tenant7.construens.com"
);
assert_eq!(
req.headers().get("x-forwarded-host").unwrap(),
"tenant7.construens.com"
);
}
}
#[cfg(all(test, feature = "handlers"))]
mod sql_grant_tests {
use super::granted_sql_databases;
fn v(items: &[&str]) -> Vec<String> {
items.iter().copied().map(String::from).collect()
}
#[test]
fn bare_sql_grants_only_the_default_database() {
assert_eq!(granted_sql_databases(&v(&["sql"]), &v(&["sql"])), v(&[""]));
assert!(granted_sql_databases(&v(&["sql"]), &v(&[])).is_empty());
assert!(granted_sql_databases(&v(&[]), &v(&["sql"])).is_empty());
}
#[test]
fn a_named_grant_is_the_intersection_and_the_site_is_the_ceiling() {
assert_eq!(
granted_sql_databases(&v(&["sql:product"]), &v(&["sql:product"])),
v(&["product"])
);
assert!(granted_sql_databases(&v(&["sql:privileged"]), &v(&["sql:product"])).is_empty());
assert_eq!(
granted_sql_databases(
&v(&["sql", "sql:product"]),
&v(&["sql", "sql:product", "sql:privileged"]),
),
v(&["", "product"])
);
}
#[test]
fn a_handler_wildcard_grants_every_name_the_site_exposes() {
assert_eq!(
granted_sql_databases(&v(&["sql:*"]), &v(&["sql:product", "sql:privileged"])),
v(&["product", "privileged"])
);
assert!(granted_sql_databases(&v(&["sql:*"]), &v(&["sql"])).is_empty());
assert!(granted_sql_databases(&v(&["sql:*"]), &v(&["sql:*"])).is_empty());
}
}
#[cfg(test)]
mod cookie_auth_tests {
use super::*;
use boatramp_core::config::CookieAuthConfig;
fn cfg(origins: &[&str]) -> CookieAuthConfig {
CookieAuthConfig {
cookie_name: "session".to_string(),
allowed_origins: origins
.iter()
.map(std::string::ToString::to_string)
.collect(),
}
}
fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
let mut h = HeaderMap::new();
for (name, value) in pairs {
h.insert(
header::HeaderName::from_bytes(name.as_bytes()).unwrap(),
HeaderValue::from_str(value).unwrap(),
);
}
h
}
#[test]
fn cookie_value_extracts_the_named_cookie() {
let h = headers(&[("cookie", "a=1; session=tok123; b=2")]);
assert_eq!(cookie_value(&h, "session").as_deref(), Some("tok123"));
assert_eq!(cookie_value(&h, "missing"), None);
assert_eq!(cookie_value(&HeaderMap::new(), "session"), None);
}
#[test]
fn referer_origin_is_the_scheme_host_port() {
assert_eq!(
referer_origin("https://app.example.com/a/b?q=1"),
Some("https://app.example.com".to_string())
);
assert_eq!(
referer_origin("http://localhost:3000/x"),
Some("http://localhost:3000".to_string())
);
assert_eq!(referer_origin("not a url"), None);
}
#[test]
fn origin_check_allows_listed_cross_origins_and_absent_signal_but_rejects_others() {
let allowed = ["https://app.example.com".to_string()];
assert!(origin_allowed(
&headers(&[("origin", "https://app.example.com")]),
&allowed
));
assert!(!origin_allowed(
&headers(&[("origin", "https://evil.example.net")]),
&allowed
));
assert!(origin_allowed(
&headers(&[("referer", "https://app.example.com/page")]),
&allowed
));
assert!(!origin_allowed(
&headers(&[("referer", "https://evil.example.net/page")]),
&allowed
));
assert!(origin_allowed(&HeaderMap::new(), &allowed));
}
#[test]
fn origin_check_auto_allows_same_origin_even_with_an_empty_allowlist() {
let empty: [String; 0] = [];
assert!(origin_allowed(
&headers(&[
("host", "app.example.com"),
("origin", "https://app.example.com"),
]),
&empty
));
assert!(origin_allowed(
&headers(&[
("host", "app.example.com"),
("referer", "https://app.example.com/dashboard"),
]),
&empty
));
assert!(origin_allowed(
&headers(&[
("host", "localhost:3000"),
("origin", "http://localhost:3000"),
]),
&empty
));
assert!(!origin_allowed(
&headers(&[
("host", "app.example.com"),
("origin", "https://evil.example.net"),
]),
&empty
));
assert!(origin_allowed(
&headers(&[
("host", "app.example.com"),
("origin", "http://app.example.com")
]),
&empty
));
}
#[test]
fn outcome_injects_a_same_origin_cookie_with_an_empty_allowlist() {
let h = headers(&[
("cookie", "session=tok"),
("host", "app.example.com"),
("origin", "https://app.example.com"),
]);
assert!(matches!(
cookie_auth_outcome(&h, Some(&cfg(&[]))),
CookieAuthOutcome::Inject(t) if t == "tok"
));
}
#[test]
fn outcome_injects_a_listed_cross_origin_cookie() {
let h = headers(&[
("cookie", "session=tok"),
("host", "api.example.com"),
("origin", "https://app.example.com"),
]);
assert!(matches!(
cookie_auth_outcome(&h, Some(&cfg(&["https://app.example.com"]))),
CookieAuthOutcome::Inject(t) if t == "tok"
));
}
#[test]
fn outcome_rejects_a_cross_origin_cookie_request() {
let h = headers(&[
("cookie", "session=tok"),
("host", "app.example.com"),
("origin", "https://evil.example.net"),
]);
assert!(matches!(
cookie_auth_outcome(&h, Some(&cfg(&["https://app.example.com"]))),
CookieAuthOutcome::Reject
));
}
#[test]
fn outcome_lets_the_authorization_header_win() {
let h = headers(&[
("cookie", "session=cookietok"),
("authorization", "Bearer headertok"),
("origin", "https://evil.example.net"), ]);
assert!(matches!(
cookie_auth_outcome(&h, Some(&cfg(&["https://app.example.com"]))),
CookieAuthOutcome::None
));
}
#[test]
fn outcome_is_none_without_a_cookie_or_config() {
assert!(matches!(
cookie_auth_outcome(&HeaderMap::new(), Some(&cfg(&["https://app.example.com"]))),
CookieAuthOutcome::None
));
assert!(matches!(
cookie_auth_outcome(&headers(&[("cookie", "session=tok")]), None),
CookieAuthOutcome::None
));
}
}