use std::collections::{BTreeMap, HashMap};
use serde_json::{Value, json};
use super::config::{
IdTokenConfig, IdentityMap, OAuth2LoginConfig, ProviderConfig, RESERVED_AUTHORIZE_PARAMS,
ReturnToConfig, 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 PROVIDER_LABEL_SINGLE: &str = "default";
const PROVIDER_LABEL_UNKNOWN: &str = "unknown";
fn provider_label(canonical_slug: &str) -> &str {
if canonical_slug.is_empty() {
PROVIDER_LABEL_SINGLE
} else {
canonical_slug
}
}
const NONCE_BYTES: usize = 32;
const MAX_RETURN_TO_BYTES: usize = 512;
const MAX_USERINFO_BYTES: usize = 65_536;
const STATE_ALG: jsonwebtoken::Algorithm = jsonwebtoken::Algorithm::HS256;
pub struct Redirect {
pub location: String,
pub set_cookie: String,
}
impl std::fmt::Debug for Redirect {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Redirect")
.field("location", &self.location)
.field("set_cookie", &"<redacted>")
.finish()
}
}
pub struct Grant {
pub metadata: Value,
pub identity: Option<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 instance_providers:
&'a std::collections::BTreeMap<String, crate::config::InstanceProviderConfig>,
pub discovery: &'a std::sync::Arc<crate::channel::oidc_discovery::DiscoveryCache>,
}
pub struct CompiledProvider {
kind: &'static str,
client_id: String,
client_secret: String,
authorize_url: String,
token_url: String,
redirect_uri: String,
client_auth: String,
scopes: Vec<String>,
extra_authorize_params: BTreeMap<String, String>,
id_token: Option<IdTokenConfig>,
id_token_verifier: Option<crate::jwt::Verifier>,
userinfo_url: Option<String>,
identity: CompiledIdentityMap,
}
impl CompiledProvider {
fn wants_oidc_nonce(&self) -> bool {
self.id_token.as_ref().is_some_and(|id| id.nonce)
}
}
struct CompiledIdentityMap {
subject: String,
login: String,
name: String,
email: String,
picture: String,
}
impl CompiledIdentityMap {
fn resolve(map: Option<&IdentityMap>) -> Self {
let pick = |v: Option<&str>, default: &str| {
v.filter(|s| !s.trim().is_empty())
.map(str::to_string)
.unwrap_or_else(|| default.to_string())
};
Self {
subject: pick(map.and_then(|m| m.subject.as_deref()), "sub"),
login: pick(map.and_then(|m| m.login.as_deref()), "preferred_username"),
name: pick(map.and_then(|m| m.name.as_deref()), "name"),
email: pick(map.and_then(|m| m.email.as_deref()), "email"),
picture: pick(map.and_then(|m| m.picture.as_deref()), "picture"),
}
}
fn extract(&self, source: &Value) -> Option<Value> {
let subject = source.get(&self.subject).and_then(value_to_string)?;
let mut identity = json!({ "subject": subject });
for (out, key) in [
("login", &self.login),
("name", &self.name),
("email", &self.email),
("picture", &self.picture),
] {
if let Some(v) = source.get(key).and_then(value_to_string) {
identity[out] = json!(v);
}
}
Some(identity)
}
}
fn value_to_string(v: &Value) -> Option<String> {
match v {
Value::String(s) if !s.is_empty() => Some(s.clone()),
Value::Number(n) => Some(n.to_string()),
Value::Bool(b) => Some(b.to_string()),
_ => None,
}
}
pub struct CompiledOAuth2Login {
channel: String,
callback_path: String,
pkce: bool,
run_workflow_on_authorize: bool,
state_cookie: StateCookieConfig,
return_to: Option<ReturnToConfig>,
state_key: jsonwebtoken::EncodingKey,
state_verifier: crate::jwt::Verifier,
providers: BTreeMap<String, CompiledProvider>,
route_selected: bool,
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("callback_path", &self.callback_path)
.field("pkce", &self.pkce)
.field("route_selected", &self.route_selected)
.field("providers", &self.providers.keys().collect::<Vec<_>>())
.finish_non_exhaustive()
}
}
impl CompiledOAuth2Login {
pub async fn compile(
cfg: &OAuth2LoginConfig,
channel: &str,
deps: &LoginDeps<'_>,
) -> Result<Self, String> {
let state_secret = resolve_secret(&cfg.state_secret, "oauth2_login.state_secret").await?;
let shared_redirect =
resolve_secret(&cfg.redirect_uri, "oauth2_login.redirect_uri").await?;
let authored = merge_instance_providers(cfg, deps.instance_providers);
let mut resolved_entries: Vec<(String, ProviderConfig)> = Vec::new();
for (slug, p) in authored {
let pfx = field_prefix(&slug);
let resolved = ProviderConfig {
kind: p.kind.clone(),
issuer: resolve_opt(&p.issuer, &format!("oauth2_login.{pfx}issuer")).await?,
authorize_url: resolve_opt(
&p.authorize_url,
&format!("oauth2_login.{pfx}authorize_url"),
)
.await?,
token_url: resolve_opt(&p.token_url, &format!("oauth2_login.{pfx}token_url"))
.await?,
client_id: resolve_opt(&p.client_id, &format!("oauth2_login.{pfx}client_id"))
.await?,
client_secret: resolve_opt(
&p.client_secret,
&format!("oauth2_login.{pfx}client_secret"),
)
.await?,
client_auth: p.client_auth.clone(),
redirect_uri: resolve_opt(
&p.redirect_uri,
&format!("oauth2_login.{pfx}redirect_uri"),
)
.await?,
scopes: p.scopes.clone(),
extra_authorize_params: p.extra_authorize_params.clone(),
id_token: p.id_token.clone(),
userinfo_url: resolve_opt(
&p.userinfo_url,
&format!("oauth2_login.{pfx}userinfo_url"),
)
.await?,
identity: p.identity.clone(),
};
resolved_entries.push((slug, resolved));
}
let resolved_cfg = assemble_resolved(cfg, shared_redirect.clone(), &resolved_entries);
validate_shape(&resolved_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 mut providers = BTreeMap::new();
for (slug, p) in &resolved_entries {
let discovered = match p.issuer.as_deref() {
Some(issuer) => Some(
deps.discovery
.resolve(issuer)
.await
.map_err(|e| format!("oauth2_login.{}issuer: {e}", field_prefix(slug)))?,
),
None => None,
};
providers.insert(
slug.clone(),
build_provider(slug, p, &shared_redirect, discovered.as_deref(), deps)?,
);
}
Ok(Self {
channel: channel.to_string(),
callback_path: cfg.callback_path.clone(),
pkce: cfg.pkce,
run_workflow_on_authorize: cfg.run_workflow_on_authorize,
state_cookie: cfg.state_cookie.clone(),
return_to: cfg.return_to.clone(),
state_key,
state_verifier,
providers,
route_selected: cfg.is_multi_provider(),
http_client: deps.http_client.clone(),
allow_private_token_urls: deps.allow_private_token_urls,
})
}
pub fn callback_path(&self) -> &str {
&self.callback_path
}
pub fn runs_workflow_on_authorize(&self) -> bool {
self.run_workflow_on_authorize
}
pub fn state_cookie_name(&self) -> &str {
&self.state_cookie.name
}
fn select(&self, slug: Option<&str>) -> Result<(&str, &CompiledProvider), OrionError> {
if self.route_selected {
let slug = slug.unwrap_or_default();
self.providers
.get_key_value(slug)
.map(|(k, p)| (k.as_str(), p))
.ok_or_else(|| {
OrionError::NotFound("no such identity provider on this channel".to_string())
})
} else {
let (k, p) = self
.providers
.iter()
.next()
.expect("a compiled block always has at least one provider");
Ok((k.as_str(), p))
}
}
pub fn require_provider(&self, slug: Option<&str>) -> Result<(), OrionError> {
self.select(slug).map(|_| ())
}
pub fn begin(
&self,
slug: Option<&str>,
contributed: Option<&Value>,
return_to: Option<&str>,
) -> Result<Redirect, OrionError> {
let (canonical_slug, provider) = match self.select(slug) {
Ok(v) => v,
Err(e) => {
if self.route_selected {
crate::metrics::record_oauth_login(
&self.channel,
PROVIDER_LABEL_UNKNOWN,
Leg::Authorize,
"unknown_provider",
);
}
return Err(e);
}
};
let nonce = random_nonce();
let oidc_nonce = provider.wants_oidc_nonce().then(random_nonce);
let verifier = self.pkce.then(random_nonce);
let mut url = url::Url::parse(&provider.authorize_url).map_err(|e| {
OrionError::internal(format!("oauth2_login authorize_url does not parse: {e}"))
})?;
{
let mut q = url.query_pairs_mut();
q.append_pair("response_type", "code");
q.append_pair("client_id", &provider.client_id);
q.append_pair("redirect_uri", &provider.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(|| provider.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 &provider.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.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);
}
if self.route_selected {
claims["provider"] = json!(canonical_slug);
}
let token = crate::jwt::sign(STATE_ALG, &self.state_key, None, &claims)
.map_err(|e| OrionError::internal(format!("could not sign the OAuth2 state: {e}")))?;
crate::metrics::record_oauth_login(
&self.channel,
provider_label(canonical_slug),
Leg::Authorize,
"ok",
);
Ok(Redirect {
location: url.into(),
set_cookie: self
.state_cookie(&token, self.state_cookie.max_age as i64)
.map_err(OrionError::internal)?,
})
}
pub async fn complete(
&self,
slug: Option<&str>,
query: &HashMap<String, String>,
jar: &[&str],
) -> Result<Grant, OrionError> {
let unknown = PROVIDER_LABEL_UNKNOWN;
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(unknown, "provider_error"));
}
let state = query
.get("state")
.ok_or_else(|| self.refuse(unknown, "state_missing"))?;
let code = query
.get("code")
.ok_or_else(|| self.refuse(unknown, "code_missing"))?;
let cookie = crate::channel::cookies::lookup(jar.iter().copied(), &self.state_cookie.name)
.ok_or_else(|| self.refuse(unknown, "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(unknown, "state_invalid")
})?;
let minted = claims
.get("nonce")
.and_then(Value::as_str)
.ok_or_else(|| self.refuse(unknown, "state_invalid"))?;
if !secret_eq(state, minted) {
return Err(self.refuse(unknown, "state_mismatch"));
}
if self.route_selected {
let route_slug = slug.unwrap_or_default();
let sealed = claims
.get("provider")
.and_then(Value::as_str)
.ok_or_else(|| self.refuse(unknown, "state_invalid"))?;
if sealed != route_slug {
return Err(self.refuse(unknown, "provider_mismatch"));
}
}
let (canonical_slug, provider) = self
.select(slug)
.map_err(|_| self.refuse(unknown, "unknown_provider"))?;
let label = provider_label(canonical_slug);
let tokens = self
.exchange(
provider,
label,
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);
}
self.stamp_provenance(&mut oauth, canonical_slug, provider.kind);
if let Some(verifier) = provider.id_token_verifier.as_ref() {
let id = provider.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(label, "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(label, "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(label, "id_token_rejected"));
}
None => {}
}
}
let identity = self
.resolve_identity(provider, canonical_slug, &oauth)
.await?;
if let Some(ref id) = identity {
oauth["identity"] = id.clone();
}
crate::metrics::record_oauth_login(&self.channel, label, Leg::Callback, "ok");
Ok(Grant {
metadata: oauth,
identity,
clear_cookie: self.state_cookie("", 0).map_err(|e| {
OrionError::internal(format!("could not clear the state cookie: {e}"))
})?,
})
}
async fn resolve_identity(
&self,
provider: &CompiledProvider,
canonical_slug: &str,
oauth: &Value,
) -> Result<Option<Value>, OrionError> {
let label = provider_label(canonical_slug);
let mut base = if let Some(claims) = oauth.get("claims") {
provider.identity.extract(claims)
} else if let Some(url) = provider.userinfo_url.as_deref() {
let access_token = oauth
.get("access_token")
.and_then(Value::as_str)
.unwrap_or_default();
let userinfo = self.fetch_userinfo(url, access_token, label).await?;
provider.identity.extract(&userinfo)
} else {
None
};
if let Some(identity) = base.as_mut() {
self.stamp_provenance(identity, canonical_slug, provider.kind);
} else if provider.id_token_verifier.is_some() || provider.userinfo_url.is_some() {
tracing::warn!(
channel = %self.channel,
"OAuth2 sign-in resolved no identity subject; check the provider's identity mapping"
);
}
Ok(base)
}
fn stamp_provenance(&self, obj: &mut Value, canonical_slug: &str, kind: &str) {
if self.route_selected
&& let Some(map) = obj.as_object_mut()
{
map.insert("provider".to_string(), json!(canonical_slug));
map.insert("kind".to_string(), json!(kind));
}
}
async fn fetch_userinfo(
&self,
url: &str,
access_token: &str,
label: &str,
) -> Result<Value, OrionError> {
if !self.allow_private_token_urls
&& let Err(msg) = crate::validation::validate_url_not_private(url).await
{
tracing::warn!(channel = %self.channel, error = %msg, "userinfo URL refused");
return Err(self.refuse(label, "userinfo_rejected"));
}
let response = self
.http_client
.get(url)
.bearer_auth(access_token)
.header(reqwest::header::ACCEPT, "application/json")
.timeout(std::time::Duration::from_secs(5))
.send()
.await
.map_err(|e| {
tracing::warn!(channel = %self.channel, error = %e, "userinfo fetch failed");
self.count(label, "userinfo_error");
OrionError::unavailable(
Unavailable::GuardBackend,
"the identity provider could not be reached",
)
})?;
if !response.status().is_success() {
tracing::warn!(
channel = %self.channel,
status = %response.status(),
"userinfo fetch rejected"
);
return Err(self.refuse(label, "userinfo_rejected"));
}
let body = crate::http_body::read_bounded(response, MAX_USERINFO_BYTES)
.await
.map_err(|e| {
tracing::warn!(channel = %self.channel, error = %e, "userinfo body");
self.refuse(label, "userinfo_rejected")
})?;
serde_json::from_slice(&body).map_err(|e| {
tracing::warn!(channel = %self.channel, error = %e, "userinfo is not JSON");
self.refuse(label, "userinfo_rejected")
})
}
async fn exchange(
&self,
provider: &CompiledProvider,
label: &str,
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", provider.redirect_uri.clone()),
];
if let Some(v) = pkce_verifier {
params.push(("code_verifier", v.to_string()));
}
let endpoint = crate::connector::oauth::TokenEndpoint {
token_url: &provider.token_url,
client_id: &provider.client_id,
client_secret: &provider.client_secret,
client_auth: &provider.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(label, "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(label, "exchange_rejected")
}
})
}
pub fn accepted_return_to(&self, query: &HashMap<String, String>) -> Option<String> {
let cfg = self.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.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, provider: &str, outcome: &'static str) -> OrionError {
self.count(provider, outcome);
OrionError::Unauthorized("sign-in could not be completed".to_string())
}
fn count(&self, provider: &str, outcome: &'static str) {
crate::metrics::record_oauth_login(&self.channel, provider, Leg::Callback, outcome);
}
}
fn field_prefix(slug: &str) -> String {
if slug.is_empty() {
String::new()
} else {
format!("providers.{slug}.")
}
}
fn effective_redirect_uri(shared: &str, slug: &str, over: Option<&str>) -> String {
over.unwrap_or(shared).replace("{provider}", slug)
}
fn effective_kind(explicit: Option<&str>, has_id_token: bool) -> &'static str {
match explicit {
Some("oidc") => "oidc",
Some("oauth2") => "oauth2",
_ if has_id_token => "oidc",
_ => "oauth2",
}
}
fn assemble_resolved(
o: &OAuth2LoginConfig,
shared_redirect: String,
entries: &[(String, ProviderConfig)],
) -> OAuth2LoginConfig {
let mut base = OAuth2LoginConfig {
kind: None,
issuer: None,
authorize_url: None,
token_url: None,
client_id: None,
client_secret: None,
client_auth: "basic".to_string(),
providers: None,
providers_from_instance: false,
redirect_uri: shared_redirect,
callback_path: o.callback_path.clone(),
scopes: Vec::new(),
extra_authorize_params: BTreeMap::new(),
pkce: o.pkce,
state_secret: o.state_secret.clone(),
state_cookie: o.state_cookie.clone(),
run_workflow_on_authorize: o.run_workflow_on_authorize,
return_to: o.return_to.clone(),
id_token: None,
userinfo_url: None,
identity: None,
};
if o.is_multi_provider() {
base.providers = Some(entries.iter().cloned().collect());
} else if let Some((_, p)) = entries.first() {
base.kind = p.kind.clone();
base.issuer = p.issuer.clone();
base.authorize_url = p.authorize_url.clone();
base.token_url = p.token_url.clone();
base.client_id = p.client_id.clone();
base.client_secret = p.client_secret.clone();
base.client_auth = p.client_auth.clone();
base.scopes = p.scopes.clone();
base.extra_authorize_params = p.extra_authorize_params.clone();
base.id_token = p.id_token.clone();
base.userinfo_url = p.userinfo_url.clone();
base.identity = p.identity.clone();
}
base
}
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",
"issuer",
"authorize_url",
"token_url",
"redirect_uri",
"userinfo_url",
];
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);
}
let bare = field.rsplit('.').next().unwrap_or(field);
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(&bare) => 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> {
validate_shared(cfg, mode)?;
let entries = cfg.provider_entries();
if cfg.is_multi_provider() && entries.is_empty() && !cfg.providers_from_instance {
return Err("oauth2_login.providers must name at least one provider".to_string());
}
for (slug, provider) in &entries {
validate_provider(cfg, slug, provider, mode)?;
}
Ok(())
}
fn merge_instance_providers(
cfg: &OAuth2LoginConfig,
instance: &BTreeMap<String, crate::config::InstanceProviderConfig>,
) -> Vec<(String, ProviderConfig)> {
let mut entries = cfg.provider_entries();
if cfg.providers_from_instance {
let declared: std::collections::HashSet<&str> =
entries.iter().map(|(s, _)| s.as_str()).collect();
let extra: Vec<(String, ProviderConfig)> = instance
.iter()
.filter(|(slug, _)| !declared.contains(slug.as_str()))
.map(|(slug, p)| (slug.clone(), ProviderConfig::from(p)))
.collect();
entries.extend(extra);
}
entries
}
fn validate_shared(cfg: &OAuth2LoginConfig, mode: ShapeCheck) -> Result<(), String> {
let multi = cfg.is_multi_provider();
if multi {
let flat_set = cfg.kind.is_some()
|| cfg.authorize_url.is_some()
|| cfg.token_url.is_some()
|| cfg.client_id.is_some()
|| cfg.client_secret.is_some()
|| cfg.id_token.is_some()
|| !cfg.scopes.is_empty()
|| !cfg.extra_authorize_params.is_empty();
if flat_set {
return Err(
"oauth2_login sets both `providers` and the flat provider fields \
(authorize_url, client_id, …); use one form or the other — the per-provider \
fields belong inside each `providers` entry"
.to_string(),
);
}
}
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());
}
validate_provider_param("callback_path", &cfg.callback_path, multi)?;
}
if !deferred(mode, "redirect_uri", &cfg.redirect_uri)? {
if multi {
if !cfg.redirect_uri.contains("{provider}") {
return Err(
"oauth2_login.redirect_uri is a template when `providers` is set and must \
contain {provider}, e.g. https://app.example.com/v1/auth/{provider}/callback"
.to_string(),
);
}
} else if cfg.redirect_uri.contains('{') {
return Err(format!(
"oauth2_login.redirect_uri '{}' carries a path parameter, but this block names a \
single provider; only a `providers` block substitutes {{provider}}",
cfg.redirect_uri
));
}
}
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)?;
}
}
}
Ok(())
}
fn validate_provider_param(field: &str, path: &str, multi: bool) -> Result<(), String> {
let params = crate::channel::routing::route_param_names(path);
if multi {
if params != ["provider"] {
return Err(format!(
"oauth2_login.{field} '{path}' must carry exactly one {{provider}} segment when \
`providers` is set, and no other path parameter"
));
}
} else if !params.is_empty() {
return Err(format!(
"oauth2_login.{field} '{path}' carries a path parameter; a single-provider callback \
is a fixed URL registered with the identity provider, so it must be static"
));
}
Ok(())
}
fn validate_provider(
cfg: &OAuth2LoginConfig,
slug: &str,
p: &ProviderConfig,
mode: ShapeCheck,
) -> Result<(), String> {
let pfx = field_prefix(slug);
if let Some(kind) = p.kind.as_deref() {
match kind {
"oidc" | "oauth2" => {}
other => {
return Err(format!(
"oauth2_login.{pfx}kind '{other}' is not supported — expected oidc or oauth2"
));
}
}
}
let discovers = p.issuer.as_deref().is_some_and(|s| !s.trim().is_empty());
if let Some(issuer) = p.issuer.as_deref() {
let field = format!("{pfx}issuer");
if !deferred(mode, &field, issuer)? {
require_https(&field, issuer)?;
}
}
for (name, value) in [
("authorize_url", &p.authorize_url),
("token_url", &p.token_url),
] {
let field = format!("{pfx}{name}");
match value.as_deref().map(str::trim).filter(|s| !s.is_empty()) {
Some(present) => {
if !deferred(mode, &field, present)? {
require_https(&field, present)?;
}
}
None if discovers => {}
None => return Err(format!("oauth2_login.{field} is required")),
}
}
for name in ["client_id", "client_secret"] {
let value = if name == "client_id" {
&p.client_id
} else {
&p.client_secret
};
if value
.as_deref()
.map(str::trim)
.filter(|s| !s.is_empty())
.is_none()
{
return Err(format!("oauth2_login.{pfx}{name} is required"));
}
}
let redirect = effective_redirect_uri(&cfg.redirect_uri, slug, p.redirect_uri.as_deref());
let redirect_field = format!("{pfx}redirect_uri");
if !deferred(mode, &redirect_field, &redirect)? {
require_https(&redirect_field, &redirect)?;
}
if !deferred(mode, &format!("{pfx}client_auth"), &p.client_auth)?
&& crate::connector::OAuth2ClientAuth::parse(&p.client_auth).is_none()
{
return Err(format!(
"oauth2_login.{pfx}client_auth '{}' is not supported — expected {}",
p.client_auth,
crate::connector::OAuth2ClientAuth::VALUES
));
}
for name in p.extra_authorize_params.keys() {
if RESERVED_AUTHORIZE_PARAMS.contains(&name.as_str()) {
return Err(format!(
"oauth2_login.{pfx}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 let Some(ref id) = p.id_token {
if !deferred(mode, &format!("{pfx}id_token.jwks_url"), &id.jwks_url)? {
crate::jwt::validate_jwks_url(&id.jwks_url)
.map_err(|e| format!("oauth2_login.{pfx}id_token.jwks_url: {e}"))?;
}
if id.issuer.is_empty() {
return Err(format!(
"oauth2_login.{pfx}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"
));
}
if id.algorithms.is_empty() {
return Err(format!(
"oauth2_login.{pfx}id_token.algorithms must not be empty"
));
}
for alg in &id.algorithms {
if !deferred(mode, &format!("{pfx}id_token.algorithms"), alg)? {
crate::jwt::parse_algorithm(alg)
.map_err(|e| format!("oauth2_login.{pfx}id_token.algorithms: {e}"))?;
}
}
}
if let Some(userinfo) = p
.userinfo_url
.as_deref()
.map(str::trim)
.filter(|s| !s.is_empty())
{
let field = format!("{pfx}userinfo_url");
if !deferred(mode, &field, userinfo)? {
require_https(&field, userinfo)?;
}
}
if let Some(ref map) = p.identity {
for (label, value) in [
("subject", &map.subject),
("login", &map.login),
("name", &map.name),
("email", &map.email),
("picture", &map.picture),
] {
if value.as_deref().is_some_and(|s| s.trim().is_empty()) {
return Err(format!(
"oauth2_login.{pfx}identity.{label} must not be empty — omit it to use the \
default claim name"
));
}
}
}
Ok(())
}
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()
))
}
pub(super) 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 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 build_provider(
slug: &str,
resolved: &ProviderConfig,
shared_redirect: &str,
discovered: Option<&crate::channel::oidc_discovery::Discovered>,
deps: &LoginDeps<'_>,
) -> Result<CompiledProvider, String> {
let pfx = field_prefix(slug);
let merge_opt = |explicit: &Option<String>, from_discovery: Option<&str>| {
explicit
.clone()
.filter(|s| !s.trim().is_empty())
.or_else(|| from_discovery.map(str::to_string))
};
let pick = |explicit: &Option<String>, from_discovery: Option<&str>, name: &str| {
merge_opt(explicit, from_discovery)
.ok_or_else(|| format!("oauth2_login.{pfx}{name} is required"))
};
let client_id = pick(&resolved.client_id, None, "client_id")?;
let client_secret = pick(&resolved.client_secret, None, "client_secret")?;
let authorize_url = pick(
&resolved.authorize_url,
discovered.map(|d| d.authorize_url.as_str()),
"authorize_url",
)?;
let token_url = pick(
&resolved.token_url,
discovered.map(|d| d.token_url.as_str()),
"token_url",
)?;
let userinfo_url = merge_opt(
&resolved.userinfo_url,
discovered.and_then(|d| d.userinfo_url.as_deref()),
);
let redirect_uri =
effective_redirect_uri(shared_redirect, slug, resolved.redirect_uri.as_deref());
let id_token: Option<IdTokenConfig> = match (&resolved.id_token, discovered) {
(Some(id), _) => Some(id.clone()),
(None, Some(d)) => Some(IdTokenConfig {
required: true,
issuer: vec![d.issuer.clone()],
audience: None,
jwks_url: d.jwks_url.clone(),
algorithms: vec!["RS256".to_string()],
nonce: true,
}),
(None, None) => None,
};
let id_token_verifier = match id_token {
Some(ref id) => Some(build_id_token_verifier(&pfx, id, &client_id, deps)?),
None => None,
};
Ok(CompiledProvider {
kind: effective_kind(resolved.kind.as_deref(), id_token.is_some()),
client_id,
client_secret,
authorize_url,
token_url,
redirect_uri,
client_auth: resolved.client_auth.clone(),
scopes: resolved.scopes.clone(),
extra_authorize_params: resolved.extra_authorize_params.clone(),
id_token,
id_token_verifier,
userinfo_url,
identity: CompiledIdentityMap::resolve(resolved.identity.as_ref()),
})
}
fn build_id_token_verifier(
pfx: &str,
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.{pfx}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)
}
async fn resolve_opt(value: &Option<String>, field: &str) -> Result<Option<String>, String> {
match value {
Some(v) => Ok(Some(resolve_secret(v, field).await?)),
None => Ok(None),
}
}
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, ProviderConfig, StateCookieConfig};
const STATE_SECRET: &str = "0123456789abcdef0123456789abcdef";
fn config() -> OAuth2LoginConfig {
OAuth2LoginConfig {
kind: None,
issuer: None,
authorize_url: Some("https://idp.example.com/authorize".to_string()),
token_url: Some("https://idp.example.com/token".to_string()),
client_id: Some("client-123".to_string()),
client_secret: Some("shhh".to_string()),
client_auth: "basic".to_string(),
providers: None,
providers_from_instance: false,
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,
userinfo_url: None,
identity: None,
}
}
fn multi_config() -> OAuth2LoginConfig {
let github = ProviderConfig {
authorize_url: Some("https://github.com/login/oauth/authorize".to_string()),
token_url: Some("https://github.com/login/oauth/access_token".to_string()),
client_id: Some("gh-client".to_string()),
client_secret: Some("gh-secret".to_string()),
client_auth: "body".to_string(),
scopes: vec!["read:user".to_string()],
..Default::default()
};
let acme = ProviderConfig {
authorize_url: Some("https://acme.example.com/authorize".to_string()),
token_url: Some("https://acme.example.com/token".to_string()),
client_id: Some("acme-client".to_string()),
client_secret: Some("acme-secret".to_string()),
client_auth: "basic".to_string(),
scopes: vec!["openid".to_string(), "profile".to_string()],
..Default::default()
};
OAuth2LoginConfig {
kind: None,
issuer: None,
authorize_url: None,
token_url: None,
client_id: None,
client_secret: None,
client_auth: "basic".to_string(),
providers: Some(BTreeMap::from([
("github".to_string(), github),
("acme".to_string(), acme),
])),
providers_from_instance: false,
redirect_uri: "https://app.example.com/v1/auth/{provider}/callback".to_string(),
callback_path: "/v1/auth/{provider}/callback".to_string(),
scopes: Vec::new(),
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,
userinfo_url: None,
identity: None,
}
}
fn no_instance() -> &'static BTreeMap<String, crate::config::InstanceProviderConfig> {
static EMPTY: std::sync::LazyLock<BTreeMap<String, crate::config::InstanceProviderConfig>> =
std::sync::LazyLock::new(BTreeMap::new);
&EMPTY
}
fn no_discovery() -> &'static std::sync::Arc<crate::channel::oidc_discovery::DiscoveryCache> {
static CACHE: std::sync::LazyLock<
std::sync::Arc<crate::channel::oidc_discovery::DiscoveryCache>,
> = std::sync::LazyLock::new(|| {
std::sync::Arc::new(crate::channel::oidc_discovery::DiscoveryCache::new(
reqwest::Client::new(),
false,
))
});
&CACHE
}
fn instance_with_iitm() -> BTreeMap<String, crate::config::InstanceProviderConfig> {
BTreeMap::from([(
"iitm".to_string(),
crate::config::InstanceProviderConfig {
authorize_url: Some("https://login.iitm.example/authorize".to_string()),
token_url: Some("https://login.iitm.example/token".to_string()),
client_id: Some("iitm-client".to_string()),
client_secret: Some("iitm-secret".to_string()),
scopes: vec!["openid".to_string()],
..Default::default()
},
)])
}
async fn compile_with_instance(
cfg: &OAuth2LoginConfig,
instance: &BTreeMap<String, crate::config::InstanceProviderConfig>,
) -> Result<CompiledOAuth2Login, String> {
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,
instance_providers: instance,
discovery: no_discovery(),
},
)
.await
}
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,
instance_providers: no_instance(),
discovery: no_discovery(),
},
)
.await
.expect("compiles")
}
async fn try_compile(cfg: &OAuth2LoginConfig) -> Result<CompiledOAuth2Login, String> {
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,
instance_providers: no_instance(),
discovery: no_discovery(),
},
)
.await
}
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()
}
fn cookie_value(set_cookie: &str) -> String {
set_cookie
.split(';')
.next()
.and_then(|p| p.split_once('='))
.map(|(_, v)| v.to_string())
.expect("a cookie value")
}
#[tokio::test]
async fn the_authorize_url_carries_what_the_rfc_requires() {
let login = compiled(&config()).await;
let redirect = login.begin(None, 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 a_slug_selects_its_provider_on_the_authorize_leg() {
let login = compiled(&multi_config()).await;
let gh = params(
&login
.begin(Some("github"), None, None)
.expect("gh")
.location,
);
assert_eq!(gh.get("client_id").map(String::as_str), Some("gh-client"));
assert_eq!(
gh.get("redirect_uri").map(String::as_str),
Some("https://app.example.com/v1/auth/github/callback")
);
let acme = params(
&login
.begin(Some("acme"), None, None)
.expect("acme")
.location,
);
assert_eq!(
acme.get("client_id").map(String::as_str),
Some("acme-client")
);
assert_eq!(
acme.get("redirect_uri").map(String::as_str),
Some("https://app.example.com/v1/auth/acme/callback")
);
assert_eq!(
acme.get("scope").map(String::as_str),
Some("openid profile")
);
}
#[tokio::test]
async fn an_unknown_slug_is_not_found() {
let login = compiled(&multi_config()).await;
let err = login.begin(Some("nope"), None, None).expect_err("must 404");
assert!(matches!(err, OrionError::NotFound(_)), "{err:?}");
assert!(login.require_provider(Some("nope")).is_err());
assert!(login.require_provider(Some("github")).is_ok());
}
#[tokio::test]
async fn the_state_seals_the_provider_and_a_mismatch_is_refused() {
let login = compiled(&multi_config()).await;
let redirect = login.begin(Some("github"), None, None).expect("a redirect");
let state = params(&redirect.location)
.get("state")
.expect("a state")
.clone();
let jar = redirect
.set_cookie
.split(';')
.next()
.expect("a cookie pair")
.to_string();
let query = HashMap::from([
("state".to_string(), state),
("code".to_string(), "whatever".to_string()),
]);
let err = login
.complete(Some("acme"), &query, &[jar.as_str()])
.await
.expect_err("must refuse");
assert!(matches!(err, OrionError::Unauthorized(_)), "{err:?}");
}
#[tokio::test]
async fn two_sign_ins_in_one_second_get_different_states() {
let login = compiled(&config()).await;
let a = login.begin(None, None, None).expect("a redirect");
let b = login.begin(None, 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, None).expect("a redirect");
let cookie = cookie_value(&forged.set_cookie);
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, 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(None, &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, 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(None, &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(None, &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(None, 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 = Some(value),
"token_url" => cfg.token_url = Some(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 = Some(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 = Some(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, 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 err = try_compile(&cfg)
.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"
);
}
#[test]
fn the_multi_provider_shape_is_checked() {
assert!(validate_shape(&multi_config(), ShapeCheck::Authoring).is_ok());
let mut cfg = multi_config();
cfg.callback_path = "/v1/auth/callback".to_string();
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("no {provider}");
assert!(err.contains("{provider}"), "{err}");
let mut cfg = multi_config();
cfg.redirect_uri = "https://app.example.com/callback".to_string();
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("static redirect");
assert!(err.contains("{provider}"), "{err}");
let mut cfg = multi_config();
cfg.client_id = Some("stray".to_string());
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("both forms");
assert!(err.contains("both"), "{err}");
let mut cfg = multi_config();
cfg.providers = Some(BTreeMap::new());
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("empty map");
assert!(err.contains("at least one"), "{err}");
}
#[test]
fn a_bad_provider_field_names_the_provider() {
let mut cfg = multi_config();
if let Some(p) = cfg.providers.as_mut().and_then(|m| m.get_mut("github")) {
p.token_url = Some("http://github.example/token".to_string());
}
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("http token_url");
assert!(err.contains("providers.github.token_url"), "{err}");
}
#[tokio::test]
async fn instance_providers_are_merged_when_opted_in() {
let mut cfg = multi_config();
cfg.providers_from_instance = true;
let login = compile_with_instance(&cfg, &instance_with_iitm())
.await
.expect("compiles");
let q = params(
&login
.begin(Some("iitm"), None, None)
.expect("iitm")
.location,
);
assert_eq!(q.get("client_id").map(String::as_str), Some("iitm-client"));
assert_eq!(
q.get("redirect_uri").map(String::as_str),
Some("https://app.example.com/v1/auth/iitm/callback")
);
assert!(login.begin(Some("github"), None, None).is_ok());
}
#[tokio::test]
async fn the_definition_wins_a_slug_clash() {
let mut cfg = multi_config(); cfg.providers_from_instance = true;
let mut instance = instance_with_iitm();
instance.insert(
"github".to_string(),
crate::config::InstanceProviderConfig {
authorize_url: Some("https://github.com/login/oauth/authorize".to_string()),
token_url: Some("https://github.com/login/oauth/access_token".to_string()),
client_id: Some("instance-gh".to_string()),
client_secret: Some("x".to_string()),
..Default::default()
},
);
let login = compile_with_instance(&cfg, &instance)
.await
.expect("compiles");
let q = params(
&login
.begin(Some("github"), None, None)
.expect("gh")
.location,
);
assert_eq!(
q.get("client_id").map(String::as_str),
Some("gh-client"),
"the definition's own github entry wins"
);
}
#[tokio::test]
async fn instance_only_with_none_supplied_is_refused() {
let mut cfg = multi_config();
cfg.providers = None;
cfg.providers_from_instance = true;
let err = compile_with_instance(&cfg, no_instance())
.await
.expect_err("no providers");
assert!(err.contains("at least one"), "{err}");
}
#[test]
fn instance_opt_in_refuses_the_flat_fields() {
let mut cfg = config(); cfg.providers_from_instance = true;
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("both forms");
assert!(err.contains("both"), "{err}");
}
#[test]
fn identity_maps_oidc_claims_by_default() {
let map = CompiledIdentityMap::resolve(None);
let claims = json!({
"sub": "abc-123",
"preferred_username": "jdoe",
"name": "J. Doe",
"email": "j@doe.example",
"picture": "https://cdn/x.png",
"extra": "ignored"
});
let id = map.extract(&claims).expect("an identity");
assert_eq!(id["subject"], "abc-123");
assert_eq!(id["login"], "jdoe");
assert_eq!(id["name"], "J. Doe");
assert_eq!(id["email"], "j@doe.example");
assert_eq!(id["picture"], "https://cdn/x.png");
}
#[test]
fn identity_map_overrides_and_coerces_the_subject() {
let map = CompiledIdentityMap::resolve(Some(&IdentityMap {
subject: Some("id".to_string()),
login: Some("login".to_string()),
picture: Some("avatar_url".to_string()),
..Default::default()
}));
let user = json!({ "id": 4210, "login": "octocat", "avatar_url": "https://gh/a.png" });
let id = map.extract(&user).expect("an identity");
assert_eq!(id["subject"], "4210", "a numeric id is stringified");
assert_eq!(id["login"], "octocat");
assert_eq!(id["picture"], "https://gh/a.png");
assert!(id.get("email").is_none());
}
#[test]
fn identity_is_none_without_a_subject() {
let map = CompiledIdentityMap::resolve(None);
assert!(map.extract(&json!({ "name": "no sub here" })).is_none());
}
#[test]
fn value_to_string_coerces_scalars_only() {
assert_eq!(value_to_string(&json!("s")).as_deref(), Some("s"));
assert_eq!(value_to_string(&json!(42)).as_deref(), Some("42"));
assert_eq!(value_to_string(&json!(true)).as_deref(), Some("true"));
assert_eq!(value_to_string(&json!("")), None);
assert_eq!(value_to_string(&json!(null)), None);
assert_eq!(value_to_string(&json!([1, 2])), None);
assert_eq!(value_to_string(&json!({"a": 1})), None);
}
#[test]
fn an_empty_identity_key_is_refused() {
let mut cfg = config();
cfg.identity = Some(IdentityMap {
subject: Some(" ".to_string()),
..Default::default()
});
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("empty key");
assert!(err.contains("identity.subject"), "{err}");
}
#[test]
fn a_plain_http_userinfo_url_is_refused() {
let mut cfg = config();
cfg.userinfo_url = Some("http://api.example/user".to_string());
let err = validate_shape(&cfg, ShapeCheck::Authoring).expect_err("http userinfo");
assert!(
err.contains("userinfo_url") && err.contains("https"),
"{err}"
);
}
#[tokio::test]
async fn a_short_state_secret_is_refused_at_compile() {
let mut cfg = config();
cfg.state_secret = "too-short".to_string();
let err = try_compile(&cfg).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 err = try_compile(&cfg).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 the_debug_rendering_does_not_carry_the_client_secret() {
let login = compiled(&config()).await;
let rendered = format!("{login:?}");
assert!(!rendered.contains("shhh"), "{rendered}");
}
}