use std::collections::HashMap;
use serde_json::{Value, json};
use super::config::{
IdTokenConfig, OAuth2LoginConfig, RESERVED_AUTHORIZE_PARAMS, StateCookieConfig,
};
use crate::errors::{OrionError, Unavailable};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Leg {
Authorize,
Callback,
}
impl Leg {
pub const fn as_str(self) -> &'static str {
match self {
Leg::Authorize => "authorize",
Leg::Callback => "callback",
}
}
}
const NONCE_BYTES: usize = 32;
const MAX_RETURN_TO_BYTES: usize = 512;
const STATE_ALG: jsonwebtoken::Algorithm = jsonwebtoken::Algorithm::HS256;
pub struct Redirect {
pub location: String,
pub set_cookie: String,
}
pub struct Grant {
pub metadata: Value,
pub clear_cookie: String,
}
impl std::fmt::Debug for Grant {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let keys: Vec<&str> = self
.metadata
.as_object()
.map(|m| m.keys().map(String::as_str).collect())
.unwrap_or_default();
f.debug_struct("Grant").field("fields", &keys).finish()
}
}
pub struct LoginDeps<'a> {
pub http_client: &'a reqwest::Client,
pub jwks: &'a std::sync::Arc<crate::jwt::jwks::JwksCache>,
pub allow_private_token_urls: bool,
}
pub struct CompiledOAuth2Login {
cfg: OAuth2LoginConfig,
channel: String,
client_id: String,
client_secret: String,
state_key: jsonwebtoken::EncodingKey,
state_verifier: crate::jwt::Verifier,
id_token_verifier: Option<crate::jwt::Verifier>,
http_client: reqwest::Client,
allow_private_token_urls: bool,
}
impl std::fmt::Debug for CompiledOAuth2Login {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CompiledOAuth2Login")
.field("channel", &self.channel)
.field("authorize_url", &self.cfg.authorize_url)
.field("callback_path", &self.cfg.callback_path)
.field("pkce", &self.cfg.pkce)
.field("oidc", &self.id_token_verifier.is_some())
.finish_non_exhaustive()
}
}
impl CompiledOAuth2Login {
pub async fn compile(
cfg: &OAuth2LoginConfig,
channel: &str,
deps: &LoginDeps<'_>,
) -> Result<Self, String> {
let client_id = resolve_secret(&cfg.client_id, "oauth2_login.client_id").await?;
let client_secret =
resolve_secret(&cfg.client_secret, "oauth2_login.client_secret").await?;
let state_secret = resolve_secret(&cfg.state_secret, "oauth2_login.state_secret").await?;
let cfg = OAuth2LoginConfig {
authorize_url: resolve_secret(&cfg.authorize_url, "oauth2_login.authorize_url").await?,
token_url: resolve_secret(&cfg.token_url, "oauth2_login.token_url").await?,
redirect_uri: resolve_secret(&cfg.redirect_uri, "oauth2_login.redirect_uri").await?,
..cfg.clone()
};
validate_shape(&cfg, ShapeCheck::Serving)?;
let state_key = crate::jwt::encoding_key(STATE_ALG, &state_secret, None)
.map_err(|e| format!("oauth2_login.state_secret: {e}"))?;
let state_decoding = crate::jwt::decoding_key(STATE_ALG, &state_secret, None)
.map_err(|e| format!("oauth2_login.state_secret: {e}"))?;
let state_verifier = crate::jwt::Verifier {
static_keys: vec![crate::jwt::StaticKey {
kid: None,
algorithm: STATE_ALG,
key: state_decoding,
}],
jwks: None,
algorithms: vec![STATE_ALG],
issuer: Vec::new(),
audience: Vec::new(),
leeway_secs: crate::jwt::DEFAULT_LEEWAY_SECS,
require_exp: true,
max_token_bytes: crate::jwt::DEFAULT_MAX_TOKEN_BYTES,
validations: std::sync::OnceLock::new(),
};
let id_token_verifier = match cfg.id_token {
Some(ref id) => Some(build_id_token_verifier(id, &client_id, deps)?),
None => None,
};
Ok(Self {
cfg,
channel: channel.to_string(),
client_id,
client_secret,
state_key,
state_verifier,
id_token_verifier,
http_client: deps.http_client.clone(),
allow_private_token_urls: deps.allow_private_token_urls,
})
}
pub fn callback_path(&self) -> &str {
&self.cfg.callback_path
}
pub fn runs_workflow_on_authorize(&self) -> bool {
self.cfg.run_workflow_on_authorize
}
pub fn state_cookie_name(&self) -> &str {
&self.cfg.state_cookie.name
}
pub fn begin(
&self,
contributed: Option<&Value>,
return_to: Option<&str>,
) -> Result<Redirect, String> {
let nonce = random_nonce();
let oidc_nonce = self.wants_oidc_nonce().then(random_nonce);
let verifier = self.cfg.pkce.then(random_nonce);
let mut url = url::Url::parse(&self.cfg.authorize_url)
.map_err(|e| format!("authorize_url does not parse: {e}"))?;
{
let mut q = url.query_pairs_mut();
q.append_pair("response_type", "code");
q.append_pair("client_id", &self.client_id);
q.append_pair("redirect_uri", &self.cfg.redirect_uri);
q.append_pair("state", &nonce);
let scopes = contributed
.and_then(|c| c.get("scopes"))
.and_then(Value::as_array)
.map(|a| {
a.iter()
.filter_map(Value::as_str)
.map(str::to_string)
.collect::<Vec<_>>()
})
.unwrap_or_else(|| self.cfg.scopes.clone());
if !scopes.is_empty() {
q.append_pair("scope", &scopes.join(" "));
}
if let Some(ref n) = oidc_nonce {
q.append_pair("nonce", n);
}
if let Some(ref v) = verifier {
q.append_pair("code_challenge", &pkce_challenge(v));
q.append_pair("code_challenge_method", "S256");
}
for (k, v) in &self.cfg.extra_authorize_params {
q.append_pair(k, v);
}
if let Some(extra) = contributed
.and_then(|c| c.get("extra_params"))
.and_then(Value::as_object)
{
for (k, v) in extra {
if RESERVED_AUTHORIZE_PARAMS.contains(&k.as_str()) {
tracing::warn!(
channel = %self.channel,
param = %k,
"Workflow tried to set an authorize parameter Orion owns; ignoring"
);
continue;
}
if let Some(v) = v.as_str() {
q.append_pair(k, v);
}
}
}
}
let now = now_secs();
let mut claims = json!({
"nonce": nonce,
"iat": now,
"exp": now + self.cfg.state_cookie.max_age,
});
if let Some(v) = verifier {
claims["pkce_verifier"] = json!(v);
}
if let Some(n) = oidc_nonce {
claims["oidc_nonce"] = json!(n);
}
if let Some(r) = return_to {
claims["return_to"] = json!(r);
}
let token = crate::jwt::sign(STATE_ALG, &self.state_key, None, &claims)
.map_err(|e| format!("could not sign the OAuth2 state: {e}"))?;
Ok(Redirect {
location: url.into(),
set_cookie: self.state_cookie(&token, self.cfg.state_cookie.max_age as i64)?,
})
}
pub async fn complete(
&self,
query: &HashMap<String, String>,
jar: &[&str],
) -> Result<Grant, OrionError> {
if let Some(err) = query.get("error") {
tracing::info!(
channel = %self.channel,
error = %err,
description = query.get("error_description").map(String::as_str).unwrap_or(""),
"OAuth2 sign-in refused at the identity provider"
);
return Err(self.refuse("provider_error"));
}
let state = query
.get("state")
.ok_or_else(|| self.refuse("state_missing"))?;
let code = query
.get("code")
.ok_or_else(|| self.refuse("code_missing"))?;
let cookie =
crate::channel::cookies::lookup(jar.iter().copied(), &self.cfg.state_cookie.name)
.ok_or_else(|| self.refuse("state_missing"))?;
let claims = self
.state_verifier
.verify(&cookie)
.await
.map_err(|reason| {
tracing::warn!(
channel = %self.channel,
reason = reason.as_str(),
"OAuth2 state cookie rejected"
);
self.refuse("state_invalid")
})?;
let minted = claims
.get("nonce")
.and_then(Value::as_str)
.ok_or_else(|| self.refuse("state_invalid"))?;
if !secret_eq(state, minted) {
return Err(self.refuse("state_mismatch"));
}
let tokens = self
.exchange(code, claims.get("pkce_verifier").and_then(Value::as_str))
.await?;
let mut oauth = json!({
"access_token": tokens.access_token,
"token_type": tokens.token_type.as_deref().unwrap_or("Bearer"),
});
for (key, value) in [
("refresh_token", tokens.refresh_token),
("id_token", tokens.id_token.clone()),
("scope", tokens.scope),
] {
if let Some(v) = value {
oauth[key] = json!(v);
}
}
if let Some(expires_in) = tokens.expires_in {
oauth["expires_in"] = json!(expires_in);
}
if let Some(return_to) = claims.get("return_to").and_then(Value::as_str) {
oauth["return_to"] = json!(return_to);
}
if let Some(verifier) = self.id_token_verifier.as_ref() {
let id = self.cfg.id_token.as_ref().expect("verifier implies config");
match tokens.id_token.as_deref() {
Some(token) => {
let verified = verifier.verify(token).await.map_err(|reason| {
tracing::warn!(
channel = %self.channel,
reason = reason.as_str(),
"OAuth2 id_token rejected"
);
self.refuse("id_token_rejected")
})?;
if id.nonce {
let minted = claims.get("oidc_nonce").and_then(Value::as_str);
let echoed = verified.get("nonce").and_then(Value::as_str);
match (minted, echoed) {
(Some(a), Some(b)) if secret_eq(a, b) => {}
_ => return Err(self.refuse("nonce_mismatch")),
}
}
oauth["claims"] = verified;
}
None if id.required => {
tracing::warn!(
channel = %self.channel,
"Token response carried no id_token, but one is required"
);
return Err(self.refuse("id_token_rejected"));
}
None => {}
}
}
crate::metrics::record_oauth_login(&self.channel, Leg::Callback, "ok");
Ok(Grant {
metadata: oauth,
clear_cookie: self.state_cookie("", 0).map_err(|e| {
OrionError::internal(format!("could not clear the state cookie: {e}"))
})?,
})
}
async fn exchange(
&self,
code: &str,
pkce_verifier: Option<&str>,
) -> Result<crate::connector::oauth::TokenResponse, OrionError> {
let mut params = vec![
("grant_type", "authorization_code".to_string()),
("code", code.to_string()),
("redirect_uri", self.cfg.redirect_uri.clone()),
];
if let Some(v) = pkce_verifier {
params.push(("code_verifier", v.to_string()));
}
let endpoint = crate::connector::oauth::TokenEndpoint {
token_url: &self.cfg.token_url,
client_id: &self.client_id,
client_secret: &self.client_secret,
client_auth: &self.cfg.client_auth,
};
crate::connector::oauth::exchange_code(
&self.http_client,
&self.channel,
endpoint,
self.allow_private_token_urls,
params,
)
.await
.map_err(|e| {
if e.retryable() {
tracing::warn!(channel = %self.channel, error = %e, "OAuth2 token exchange failed");
self.count("exchange_error");
OrionError::unavailable(
Unavailable::GuardBackend,
"the identity provider could not be reached",
)
} else {
tracing::warn!(channel = %self.channel, error = %e, "OAuth2 token exchange rejected");
self.refuse("exchange_rejected")
}
})
}
fn wants_oidc_nonce(&self) -> bool {
self.cfg.id_token.as_ref().is_some_and(|id| id.nonce)
}
pub fn accepted_return_to(&self, query: &HashMap<String, String>) -> Option<String> {
let cfg = self.cfg.return_to.as_ref()?;
let value = query.get(&cfg.param)?;
if value.len() > MAX_RETURN_TO_BYTES {
return None;
}
let candidate = url::Url::parse(value).ok()?;
cfg.allow_list
.iter()
.any(|entry| permits_return_to(entry, &candidate))
.then(|| value.clone())
}
fn state_cookie(&self, value: &str, max_age: i64) -> Result<String, String> {
let StateCookieConfig {
ref name,
secure,
ref same_site,
ref path,
..
} = self.cfg.state_cookie;
super::cookies::format_set_cookie(&json!({
"name": name,
"value": value,
"path": path,
"max_age": max_age,
"same_site": same_site,
"secure": secure,
"http_only": true,
}))
}
fn refuse(&self, outcome: &'static str) -> OrionError {
self.count(outcome);
OrionError::Unauthorized("sign-in could not be completed".to_string())
}
fn count(&self, outcome: &'static str) {
crate::metrics::record_oauth_login(&self.channel, Leg::Callback, outcome);
}
}
const MAX_STATE_COOKIE_MAX_AGE_SECS: u64 = 86_400;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ShapeCheck {
Authoring,
Serving,
}
pub const SECRET_RESOLVED_FIELDS: &[&str] = &[
"client_id",
"client_secret",
"state_secret",
"authorize_url",
"token_url",
"redirect_uri",
];
fn deferred(mode: ShapeCheck, field: &str, value: &str) -> Result<bool, String> {
let is_var = value.starts_with(crate::config::vars::VAR_SCHEME);
if !is_var && !crate::connector::secrets::is_resolvable_reference(value) {
return Ok(false);
}
match mode {
ShapeCheck::Serving => Err(format!(
"oauth2_login.{field} still holds '{value}' after resolution; nothing resolves a \
reference in this field, so its text would reach the identity provider"
)),
ShapeCheck::Authoring if is_var || SECRET_RESOLVED_FIELDS.contains(&field) => Ok(true),
ShapeCheck::Authoring => Err(format!(
"oauth2_login.{field} holds '{value}', but a secret reference is resolved only in \
{}; for a per-environment value here use var://name",
SECRET_RESOLVED_FIELDS.join(", ")
)),
}
}
pub fn validate_shape(cfg: &OAuth2LoginConfig, mode: ShapeCheck) -> Result<(), String> {
for (field, value) in [
("authorize_url", &cfg.authorize_url),
("token_url", &cfg.token_url),
("redirect_uri", &cfg.redirect_uri),
] {
if !deferred(mode, field, value)? {
require_https(field, value)?;
}
}
if !deferred(mode, "callback_path", &cfg.callback_path)? {
if cfg.callback_path.trim().is_empty() || !cfg.callback_path.starts_with('/') {
return Err("oauth2_login.callback_path must be an absolute path, e.g. \
/v1/auth/github/callback"
.to_string());
}
if cfg.callback_path.contains('{') {
return Err(format!(
"oauth2_login.callback_path '{}' carries a path parameter; the callback is a \
fixed URL registered with the identity provider, so it must be static",
cfg.callback_path
));
}
}
if !deferred(mode, "client_auth", &cfg.client_auth)?
&& crate::connector::OAuth2ClientAuth::parse(&cfg.client_auth).is_none()
{
return Err(format!(
"oauth2_login.client_auth '{}' is not supported — expected {}",
cfg.client_auth,
crate::connector::OAuth2ClientAuth::VALUES
));
}
for name in cfg.extra_authorize_params.keys() {
if RESERVED_AUTHORIZE_PARAMS.contains(&name.as_str()) {
return Err(format!(
"oauth2_login.extra_authorize_params sets '{name}', which Orion owns. \
Overriding it would disable the protection it carries — the reserved \
set is: {}",
RESERVED_AUTHORIZE_PARAMS.join(", ")
));
}
}
if cfg.state_cookie.max_age == 0 {
return Err(
"oauth2_login.state_cookie.max_age must be greater than zero — it is \
also the state token's expiry"
.to_string(),
);
}
if cfg.state_cookie.max_age > MAX_STATE_COOKIE_MAX_AGE_SECS {
return Err(format!(
"oauth2_login.state_cookie.max_age is {} seconds, above the {MAX_STATE_COOKIE_MAX_AGE_SECS} \
second ceiling ({} days). It is the window to finish one consent screen, not a \
session lifetime, and a long one keeps a replayable state token valid for as long \
as it lasts",
cfg.state_cookie.max_age,
MAX_STATE_COOKIE_MAX_AGE_SECS / 86_400,
));
}
if !deferred(mode, "state_cookie.same_site", &cfg.state_cookie.same_site)? {
match cfg.state_cookie.same_site.to_ascii_lowercase().as_str() {
"lax" | "none" => {}
"strict" => {
return Err(
"oauth2_login.state_cookie.same_site = \"strict\" would withhold the \
cookie on the callback, which is a top-level cross-site GET from the \
identity provider — every sign-in would fail the state check. Use \
\"lax\"."
.to_string(),
);
}
_ => {
return Err(format!(
"oauth2_login.state_cookie.same_site '{}' is not valid — Lax or None",
cfg.state_cookie.same_site
));
}
}
}
if let Some(ref rt) = cfg.return_to {
if rt.param.trim().is_empty() {
return Err("oauth2_login.return_to.param must not be empty".to_string());
}
if rt.allow_list.is_empty() {
return Err(
"oauth2_login.return_to.allow_list must list at least one permitted \
destination prefix — an empty list accepts nothing, so omit the \
whole block instead"
.to_string(),
);
}
for prefix in &rt.allow_list {
if !deferred(mode, "return_to.allow_list", prefix)? {
require_https("return_to.allow_list entry", prefix)?;
}
}
}
if let Some(ref id) = cfg.id_token {
if !deferred(mode, "id_token.jwks_url", &id.jwks_url)? {
crate::jwt::validate_jwks_url(&id.jwks_url)
.map_err(|e| format!("oauth2_login.id_token.jwks_url: {e}"))?;
}
if id.issuer.is_empty() {
return Err(
"oauth2_login.id_token.issuer must list at least one accepted \
issuer — an unchecked `iss` accepts a token from any provider \
whose key happens to be in the JWKS"
.to_string(),
);
}
if id.algorithms.is_empty() {
return Err("oauth2_login.id_token.algorithms must not be empty".to_string());
}
for alg in &id.algorithms {
if !deferred(mode, "id_token.algorithms", alg)? {
crate::jwt::parse_algorithm(alg)
.map_err(|e| format!("oauth2_login.id_token.algorithms: {e}"))?;
}
}
}
Ok(())
}
fn permits_return_to(entry: &str, candidate: &url::Url) -> bool {
let Ok(allowed) = url::Url::parse(entry) else {
return false;
};
if allowed.origin() != candidate.origin() {
return false;
}
let (allowed_path, candidate_path) = (allowed.path(), candidate.path());
if let Some(base) = allowed_path.strip_suffix('/') {
return candidate_path == base || candidate_path.starts_with(allowed_path);
}
candidate_path == allowed_path || candidate_path.starts_with(&format!("{allowed_path}/"))
}
fn require_https(field: &str, value: &str) -> Result<(), String> {
let url = url::Url::parse(value)
.map_err(|e| format!("oauth2_login.{field} '{value}' is not a URL: {e}"))?;
if url.scheme() == "https" {
return Ok(());
}
if url.scheme() == "http" && is_loopback_host(&url) {
return Ok(());
}
Err(format!(
"oauth2_login.{field} must be https — '{value}' is {}. The client secret, the \
authorization code and the session that follows all travel over it. Plain http \
is accepted only on a loopback host (localhost, 127.0.0.1, [::1]) for local \
development",
url.scheme()
))
}
fn is_loopback_host(url: &url::Url) -> bool {
match url.host() {
Some(url::Host::Ipv4(ip)) => ip.is_loopback(),
Some(url::Host::Ipv6(ip)) => ip.is_loopback(),
Some(url::Host::Domain(host)) => host == "localhost" || host.ends_with(".localhost"),
None => false,
}
}
fn build_id_token_verifier(
cfg: &IdTokenConfig,
client_id: &str,
deps: &LoginDeps<'_>,
) -> Result<crate::jwt::Verifier, String> {
let algorithms = cfg
.algorithms
.iter()
.map(|a| crate::jwt::parse_algorithm(a))
.collect::<Result<Vec<_>, _>>()
.map_err(|e| format!("oauth2_login.id_token.algorithms: {e}"))?;
Ok(crate::jwt::Verifier {
static_keys: Vec::new(),
jwks: Some(crate::jwt::JwksSource {
url: cfg.jwks_url.clone(),
cache: std::sync::Arc::clone(deps.jwks),
}),
algorithms,
issuer: cfg.issuer.clone(),
audience: cfg
.audience
.clone()
.unwrap_or_else(|| vec![client_id.to_string()]),
leeway_secs: crate::jwt::DEFAULT_LEEWAY_SECS,
require_exp: true,
max_token_bytes: crate::jwt::DEFAULT_MAX_TOKEN_BYTES,
validations: std::sync::OnceLock::new(),
})
}
async fn resolve_secret(value: &str, field: &str) -> Result<String, String> {
let resolved = crate::connector::secrets::resolve_secret_string(value, field).await?;
if resolved.is_empty() {
return Err(format!("{field} resolved to an empty value"));
}
Ok(resolved)
}
fn random_nonce() -> String {
crate::crypto::encode_bytes(
crate::crypto::Codec::Base64Url,
&crate::crypto::random_bytes(NONCE_BYTES),
)
}
fn pkce_challenge(verifier: &str) -> String {
use sha2::Digest as _;
crate::crypto::encode_bytes(
crate::crypto::Codec::Base64Url,
&sha2::Sha256::digest(verifier.as_bytes()),
)
}
fn secret_eq(a: &str, b: &str) -> bool {
use sha2::Digest as _;
let da: [u8; 32] = sha2::Sha256::digest(a.as_bytes()).into();
let db: [u8; 32] = sha2::Sha256::digest(b.as_bytes()).into();
crate::config::constant_time_eq(&da, &db)
}
fn now_secs() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::channel::config::{OAuth2LoginConfig, StateCookieConfig};
const STATE_SECRET: &str = "0123456789abcdef0123456789abcdef";
fn config() -> OAuth2LoginConfig {
OAuth2LoginConfig {
authorize_url: "https://idp.example.com/authorize".to_string(),
token_url: "https://idp.example.com/token".to_string(),
client_id: "client-123".to_string(),
client_secret: "shhh".to_string(),
client_auth: "basic".to_string(),
redirect_uri: "https://app.example.com/v1/auth/idp/callback".to_string(),
callback_path: "/v1/auth/idp/callback".to_string(),
scopes: vec!["read:user".to_string()],
extra_authorize_params: Default::default(),
pkce: true,
state_secret: STATE_SECRET.to_string(),
state_cookie: StateCookieConfig::default(),
run_workflow_on_authorize: false,
return_to: None,
id_token: None,
}
}
async fn compiled(cfg: &OAuth2LoginConfig) -> CompiledOAuth2Login {
let jwks = std::sync::Arc::new(crate::jwt::jwks::JwksCache::new(
reqwest::Client::new(),
false,
));
CompiledOAuth2Login::compile(
cfg,
"signin",
&LoginDeps {
http_client: &reqwest::Client::new(),
jwks: &jwks,
allow_private_token_urls: false,
},
)
.await
.expect("compiles")
}
fn params(location: &str) -> HashMap<String, String> {
url::Url::parse(location)
.expect("a URL")
.query_pairs()
.map(|(k, v)| (k.into_owned(), v.into_owned()))
.collect()
}
#[tokio::test]
async fn the_authorize_url_carries_what_the_rfc_requires() {
let login = compiled(&config()).await;
let redirect = login.begin(None, None).expect("a redirect");
let q = params(&redirect.location);
assert_eq!(q.get("response_type").map(String::as_str), Some("code"));
assert_eq!(q.get("client_id").map(String::as_str), Some("client-123"));
assert_eq!(
q.get("redirect_uri").map(String::as_str),
Some("https://app.example.com/v1/auth/idp/callback")
);
assert_eq!(q.get("scope").map(String::as_str), Some("read:user"));
assert!(q.contains_key("state"));
assert_eq!(
q.get("code_challenge_method").map(String::as_str),
Some("S256")
);
assert!(redirect.set_cookie.contains("HttpOnly"));
assert!(redirect.set_cookie.contains("Secure"));
assert!(redirect.set_cookie.contains("SameSite=Lax"));
assert!(redirect.set_cookie.contains("Max-Age=600"));
}
#[tokio::test]
async fn two_sign_ins_in_one_second_get_different_states() {
let login = compiled(&config()).await;
let a = login.begin(None, None).expect("a redirect");
let b = login.begin(None, None).expect("a redirect");
assert_ne!(
params(&a.location).get("state"),
params(&b.location).get("state")
);
assert_ne!(
params(&a.location).get("code_challenge"),
params(&b.location).get("code_challenge")
);
assert_ne!(a.set_cookie, b.set_cookie);
}
#[test]
fn the_pkce_challenge_matches_the_rfc_vector() {
assert_eq!(
pkce_challenge("dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"),
"E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
);
}
#[tokio::test]
async fn a_state_token_from_another_key_is_rejected() {
let login = compiled(&config()).await;
let mut other = config();
other.state_secret = "fedcba9876543210fedcba9876543210".to_string();
let attacker = compiled(&other).await;
let forged = attacker.begin(None, None).expect("a redirect");
let cookie = forged
.set_cookie
.split(';')
.next()
.and_then(|p| p.split_once('='))
.map(|(_, v)| v.to_string())
.expect("a cookie value");
let claims = login.state_verifier.verify(&cookie).await;
assert!(
claims.is_err(),
"a state signed with another key must not verify"
);
}
#[tokio::test]
async fn a_callback_without_the_cookie_is_refused_before_the_exchange() {
let login = compiled(&config()).await;
let redirect = login.begin(None, None).expect("a redirect");
let state = params(&redirect.location)
.get("state")
.expect("a state")
.clone();
let query = HashMap::from([
("state".to_string(), state),
("code".to_string(), "whatever".to_string()),
]);
let err = login.complete(&query, &[]).await.expect_err("must refuse");
assert!(
matches!(err, OrionError::Unauthorized(_)),
"expected a 401, got {err:?}"
);
}
#[tokio::test]
async fn a_state_that_does_not_match_the_cookie_is_refused() {
let login = compiled(&config()).await;
let redirect = login.begin(None, None).expect("a redirect");
let jar = redirect
.set_cookie
.split(';')
.next()
.expect("a cookie pair")
.to_string();
let query = HashMap::from([
("state".to_string(), "not-the-minted-one".to_string()),
("code".to_string(), "whatever".to_string()),
]);
let err = login
.complete(&query, &[jar.as_str()])
.await
.expect_err("must refuse");
assert!(matches!(err, OrionError::Unauthorized(_)), "{err:?}");
}
#[tokio::test]
async fn a_provider_error_is_refused_without_looking_at_the_state() {
let login = compiled(&config()).await;
let query = HashMap::from([("error".to_string(), "access_denied".to_string())]);
let err = login.complete(&query, &[]).await.expect_err("must refuse");
assert!(matches!(err, OrionError::Unauthorized(_)), "{err:?}");
}
#[test]
fn a_reserved_authorize_parameter_is_refused_at_the_door() {
let mut cfg = config();
cfg.extra_authorize_params
.insert("state".to_string(), "attacker-chosen".to_string());
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("must refuse");
assert!(err.contains("state"), "{err}");
}
#[tokio::test]
async fn a_workflow_cannot_contribute_a_reserved_parameter() {
let mut cfg = config();
cfg.run_workflow_on_authorize = true;
let login = compiled(&cfg).await;
let contributed = json!({
"extra_params": { "state": "attacker-chosen", "login_hint": "a@b.com" }
});
let redirect = login.begin(Some(&contributed), None).expect("a redirect");
let q = params(&redirect.location);
assert_eq!(q.get("login_hint").map(String::as_str), Some("a@b.com"));
assert_ne!(q.get("state").map(String::as_str), Some("attacker-chosen"));
}
#[test]
fn http_endpoints_are_refused() {
for field in ["authorize_url", "token_url", "redirect_uri"] {
let mut cfg = config();
let value = "http://idp.example.com/x".to_string();
match field {
"authorize_url" => cfg.authorize_url = value,
"token_url" => cfg.token_url = value,
_ => cfg.redirect_uri = value,
}
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err(field);
assert!(err.contains("https"), "{field}: {err}");
}
}
#[test]
fn plain_http_is_accepted_only_on_loopback() {
for host in ["localhost", "127.0.0.1", "[::1]", "app.localhost"] {
let mut cfg = config();
cfg.token_url = format!("http://{host}:8080/token");
assert!(
validate_shape(&cfg, ShapeCheck::Authoring).is_ok(),
"{host} should be accepted"
);
}
for host in ["localhost.evil.test", "127.0.0.1.evil.test", "10.0.0.1"] {
let mut cfg = config();
cfg.token_url = format!("http://{host}/token");
assert!(
validate_shape(&cfg, ShapeCheck::Authoring).is_err(),
"{host} should be refused"
);
}
}
#[test]
fn a_reference_is_deferred_at_authoring_and_refused_when_serving() {
for value in [
"var://redirect",
"env://OAUTH_REDIRECT_URI",
"vault://kv/app#redirect",
] {
let mut cfg = config();
cfg.redirect_uri = value.to_string();
assert!(
validate_shape(&cfg, ShapeCheck::Authoring).is_ok(),
"{value} is deferred at authoring"
);
let err = validate_shape(&cfg, ShapeCheck::Serving).expect_err(value);
assert!(err.contains("redirect_uri"), "{value}: {err}");
}
let mut cfg = config();
cfg.callback_path = "var://callback".to_string();
cfg.client_auth = "var://client_auth".to_string();
cfg.state_cookie.same_site = "var://same_site".to_string();
assert!(validate_shape(&cfg, ShapeCheck::Authoring).is_ok());
let mut cfg = config();
cfg.callback_path = "env://CALLBACK_PATH".to_string();
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("nothing resolves it");
assert!(
err.contains("callback_path") && err.contains("var://"),
"{err}"
);
}
#[tokio::test]
async fn compile_resolves_the_redirect_uri_and_checks_the_result() {
unsafe {
std::env::set_var(
"ORION_TEST_OAUTH2_UNIT_REDIRECT_HTTPS",
"https://app.example.com/v1/auth/idp/callback",
);
std::env::set_var(
"ORION_TEST_OAUTH2_UNIT_REDIRECT_HTTP",
"http://app.example.com/v1/auth/idp/callback",
);
}
let mut cfg = config();
cfg.redirect_uri = "env://ORION_TEST_OAUTH2_UNIT_REDIRECT_HTTPS".to_string();
let login = compiled(&cfg).await;
let redirect = login.begin(None, None).expect("a redirect");
let q = params(&redirect.location);
assert_eq!(
q.get("redirect_uri").map(String::as_str),
Some("https://app.example.com/v1/auth/idp/callback")
);
let mut cfg = config();
cfg.redirect_uri = "env://ORION_TEST_OAUTH2_UNIT_REDIRECT_HTTP".to_string();
let jwks = std::sync::Arc::new(crate::jwt::jwks::JwksCache::new(
reqwest::Client::new(),
false,
));
let err = CompiledOAuth2Login::compile(
&cfg,
"signin",
&LoginDeps {
http_client: &reqwest::Client::new(),
jwks: &jwks,
allow_private_token_urls: false,
},
)
.await
.expect_err("plain http after resolution");
assert!(err.contains("https"), "{err}");
}
#[test]
fn a_strict_state_cookie_is_refused_with_the_reason() {
let mut cfg = config();
cfg.state_cookie.same_site = "strict".to_string();
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("must refuse");
assert!(err.contains("cross-site"), "{err}");
}
#[test]
fn a_parameterised_or_self_referencing_callback_is_refused() {
let mut cfg = config();
cfg.callback_path = "/v1/auth/{provider}/callback".to_string();
assert!(validate_shape(&cfg, ShapeCheck::Authoring).is_err());
let mut cfg = config();
cfg.callback_path = "v1/auth/idp/callback".to_string();
assert!(
validate_shape(&cfg, ShapeCheck::Authoring).is_err(),
"must be absolute"
);
}
#[tokio::test]
async fn a_short_state_secret_is_refused_at_compile() {
let mut cfg = config();
cfg.state_secret = "too-short".to_string();
let jwks = std::sync::Arc::new(crate::jwt::jwks::JwksCache::new(
reqwest::Client::new(),
false,
));
let err = CompiledOAuth2Login::compile(
&cfg,
"signin",
&LoginDeps {
http_client: &reqwest::Client::new(),
jwks: &jwks,
allow_private_token_urls: false,
},
)
.await
.expect_err("must refuse");
assert!(err.contains("RFC 7518"), "{err}");
}
#[tokio::test]
async fn state_cookie_max_age_is_bounded_at_both_ends() {
let mut cfg = config();
cfg.state_cookie.max_age = 0;
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("zero must be refused");
assert!(err.contains("greater than zero"), "{err}");
for absurd in [u64::MAX, u64::MAX / 2, MAX_STATE_COOKIE_MAX_AGE_SECS + 1] {
cfg.state_cookie.max_age = absurd;
let err =
validate_shape(&cfg, ShapeCheck::Authoring).expect_err("{absurd} must be refused");
assert!(err.contains("max_age"), "{err}");
assert!(err.contains("ceiling"), "{err}");
}
for ok in [1, 600, MAX_STATE_COOKIE_MAX_AGE_SECS] {
cfg.state_cookie.max_age = ok;
assert!(
validate_shape(&cfg, ShapeCheck::Authoring).is_ok(),
"{ok} should be accepted"
);
}
}
#[tokio::test]
async fn an_out_of_range_max_age_is_refused_at_compile() {
let mut cfg = config();
cfg.state_cookie.max_age = u64::MAX;
let jwks = std::sync::Arc::new(crate::jwt::jwks::JwksCache::new(
reqwest::Client::new(),
false,
));
let err = CompiledOAuth2Login::compile(
&cfg,
"signin",
&LoginDeps {
http_client: &reqwest::Client::new(),
jwks: &jwks,
allow_private_token_urls: false,
},
)
.await
.expect_err("must refuse");
assert!(err.contains("max_age"), "{err}");
}
#[tokio::test]
async fn return_to_is_filtered_against_the_allow_list() {
let mut cfg = config();
cfg.return_to = Some(crate::channel::ReturnToConfig {
param: "next".to_string(),
allow_list: vec!["https://app.example.com".to_string()],
});
let login = compiled(&cfg).await;
for value in [
"https://app.example.com/dashboard",
"https://app.example.com/",
"https://app.example.com",
"https://app.example.com/a/b?q=1#frag",
] {
let permitted = HashMap::from([("next".to_string(), value.to_string())]);
assert_eq!(
login.accepted_return_to(&permitted).as_deref(),
Some(value),
"{value}"
);
}
for value in [
"https://evil.example.com/",
"https://app.example.com.evil.test/steal",
"https://app.example.com.evil.test",
"https://app.example.com@evil.test/steal",
"http://app.example.com/dashboard",
"https://app.example.com:8443/dashboard",
"/dashboard",
"javascript:alert(1)",
"",
] {
let refused = HashMap::from([("next".to_string(), value.to_string())]);
assert_eq!(login.accepted_return_to(&refused), None, "{value}");
}
}
#[tokio::test]
async fn return_to_path_matching_cuts_at_a_segment_boundary() {
let mut cfg = config();
cfg.return_to = Some(crate::channel::ReturnToConfig {
param: "next".to_string(),
allow_list: vec!["https://app.example.com/app".to_string()],
});
let login = compiled(&cfg).await;
for value in [
"https://app.example.com/app",
"https://app.example.com/app/",
"https://app.example.com/app/home",
] {
let permitted = HashMap::from([("next".to_string(), value.to_string())]);
assert_eq!(
login.accepted_return_to(&permitted).as_deref(),
Some(value),
"{value}"
);
}
for value in [
"https://app.example.com/application",
"https://app.example.com/appliance/x",
"https://app.example.com/other",
"https://app.example.com/",
] {
let refused = HashMap::from([("next".to_string(), value.to_string())]);
assert_eq!(login.accepted_return_to(&refused), None, "{value}");
}
}
#[tokio::test]
async fn return_to_entry_means_the_same_with_or_without_a_trailing_slash() {
for entry in [
"https://app.example.com/app",
"https://app.example.com/app/",
] {
let mut cfg = config();
cfg.return_to = Some(crate::channel::ReturnToConfig {
param: "next".to_string(),
allow_list: vec![entry.to_string()],
});
let login = compiled(&cfg).await;
for (value, admitted) in [
("https://app.example.com/app", true),
("https://app.example.com/app/home", true),
("https://app.example.com/application", false),
] {
let q = HashMap::from([("next".to_string(), value.to_string())]);
assert_eq!(
login.accepted_return_to(&q).is_some(),
admitted,
"entry {entry} / value {value}"
);
}
}
}
#[tokio::test]
async fn the_debug_rendering_does_not_carry_the_client_secret() {
let login = compiled(&config()).await;
let rendered = format!("{login:?}");
assert!(!rendered.contains("shhh"), "{rendered}");
}
}