use crate::oauth2::{providers, OAuth2Provider, OAuthError};
pub use crate::oauth2::{open_flow, seal_flow, NormalizedUser, OAuth2Flow};
pub const SSO_FLOW_COOKIE: &str = "rustango_admin_sso_flow";
#[derive(Debug, Clone)]
pub struct ResolvedSso {
pub provider: String,
pub issuer_url: Option<String>,
pub client_id: String,
pub client_secret: String,
pub redirect_uri: String,
pub scopes: Option<Vec<String>>,
}
#[derive(Debug)]
pub enum SsoError {
UnknownProvider(String),
MissingIssuer,
NotEnabled,
EmailNotVerified,
NoMatchingUser(String),
Inactive,
Secret(String),
Config(String),
Oauth(OAuthError),
}
impl From<OAuthError> for SsoError {
fn from(e: OAuthError) -> Self {
SsoError::Oauth(e)
}
}
impl std::fmt::Display for SsoError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
SsoError::UnknownProvider(p) => write!(f, "unknown SSO provider: {p}"),
SsoError::MissingIssuer => write!(f, "provider=oidc requires an issuer_url"),
SsoError::NotEnabled => write!(f, "SSO is not enabled"),
SsoError::EmailNotVerified => write!(f, "the IdP email is not verified"),
SsoError::NoMatchingUser(e) => write!(f, "no admin account for {e}"),
SsoError::Inactive => write!(f, "the matched account is inactive"),
SsoError::Secret(m) => write!(f, "could not resolve the SSO client secret: {m}"),
SsoError::Config(m) => write!(f, "SSO misconfigured: {m}"),
SsoError::Oauth(e) => write!(f, "SSO handshake failed: {e}"),
}
}
}
impl std::error::Error for SsoError {}
pub async fn build_provider(cfg: &ResolvedSso) -> Result<OAuth2Provider, SsoError> {
if cfg.client_id.trim().is_empty() {
return Err(SsoError::Config("client_id is empty".into()));
}
let (id, secret, redirect) = (
cfg.client_id.clone(),
cfg.client_secret.clone(),
cfg.redirect_uri.clone(),
);
let provider = match cfg.provider.as_str() {
"google" => providers::google(id, secret, redirect),
"microsoft" => providers::microsoft(id, secret, redirect),
"github" => providers::github(id, secret, redirect),
"gitlab" => providers::gitlab(id, secret, redirect),
"discord" => providers::discord(id, secret, redirect),
"oidc" => {
let issuer = cfg.issuer_url.as_deref().ok_or(SsoError::MissingIssuer)?;
OAuth2Provider::from_discovery("oidc", issuer, id, secret, redirect).await?
}
other => return Err(SsoError::UnknownProvider(other.to_owned())),
};
let provider = match &cfg.scopes {
Some(s) if !s.is_empty() => provider.with_scopes(s.iter().cloned()),
_ => provider,
};
Ok(provider)
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct ProviderButton {
pub slug: String,
pub label: String,
pub login_url: String,
}
#[must_use]
pub fn parse_scopes(scopes: Option<&str>) -> Option<Vec<String>> {
let s = scopes?.trim();
if s.is_empty() {
return None;
}
Some(s.split_whitespace().map(str::to_owned).collect())
}
pub fn resolve_secret_ref_env(reference: &str) -> Result<String, SsoError> {
if let Some(var) = reference.strip_prefix("env://") {
std::env::var(var).map_err(|_| SsoError::Secret(format!("env var `{var}` is unset")))
} else {
Ok(reference.to_owned())
}
}
pub fn verified_email(user: &NormalizedUser) -> Result<&str, SsoError> {
match (&user.email, user.email_verified) {
(Some(e), true) if !e.is_empty() => Ok(e.as_str()),
_ => Err(SsoError::EmailNotVerified),
}
}
use axum::{
extract::{Path, Query, State},
http::{header, HeaderMap, HeaderValue},
response::{IntoResponse, Redirect, Response},
routing::get,
Router,
};
use super::session::{self, AdminSession, SESSION_COOKIE};
use super::urls::AppState;
use super::user::AdminUser;
use crate::core::Model as _;
#[derive(serde::Deserialize)]
struct CallbackParams {
code: Option<String>,
state: Option<String>,
error: Option<String>,
}
pub(crate) fn sso_router(state: AppState) -> Router {
Router::new()
.route("/login/sso/{slug}", get(sso_begin))
.route("/login/sso/{slug}/callback", get(sso_callback))
.with_state(state)
}
fn login_path(state: &AppState) -> String {
let p = &state.config.admin_prefix;
if p.is_empty() {
"/login".to_owned()
} else {
format!("{p}/login")
}
}
fn derive_bare_redirect(headers: &HeaderMap, state: &AppState, slug: &str) -> Option<String> {
let host = headers.get(header::HOST)?.to_str().ok()?;
let scheme = headers
.get("x-forwarded-proto")
.and_then(|v| v.to_str().ok())
.map(|s| s.split(',').next().unwrap_or(s).trim())
.filter(|s| !s.is_empty())
.unwrap_or("https");
Some(format!(
"{scheme}://{host}{}/sso/{slug}/callback",
login_path(state)
))
}
fn login_error(state: &AppState, code: &str) -> Response {
Redirect::to(&format!("{}?sso_error={code}", login_path(state))).into_response()
}
fn cookie_attrs(secure: bool) -> &'static str {
if secure {
"; Secure"
} else {
""
}
}
fn read_cookie(headers: &HeaderMap, name: &str) -> Option<String> {
let raw = headers.get(header::COOKIE)?.to_str().ok()?;
raw.split(';')
.filter_map(|kv| kv.trim().split_once('='))
.find(|(k, _)| *k == name)
.map(|(_, v)| v.to_owned())
}
async fn sso_begin(
State(state): State<AppState>,
Path(slug): Path<String>,
headers: HeaderMap,
) -> Response {
let Some(secret) = state.config.session_secret.as_ref() else {
return login_error(&state, "disabled");
};
let Some(redirect_uri) = derive_bare_redirect(&headers, &state, &slug) else {
return login_error(&state, "config");
};
let cfg = match super::sso_provider::resolve_by_slug(&state.pool, &slug, redirect_uri).await {
Ok(Some(c)) => c,
Ok(None) => return login_error(&state, "disabled"),
Err(e) => {
tracing::error!(target: "rustango::admin::sso", "begin resolve: {e}");
return login_error(&state, "config");
}
};
let provider = match build_provider(&cfg).await {
Ok(p) => p,
Err(e) => {
tracing::error!(target: "rustango::admin::sso", "begin: {e}");
return login_error(&state, "config");
}
};
let (url, flow) = provider.begin();
let sealed = seal_flow(&flow, secret.key());
let cookie = format!(
"{SSO_FLOW_COOKIE}={sealed}; Path=/; HttpOnly; SameSite=Lax; Max-Age=600{s}",
s = cookie_attrs(state.config.secure_cookies),
);
let mut resp = Redirect::to(&url).into_response();
if let Ok(v) = HeaderValue::from_str(&cookie) {
resp.headers_mut().insert(header::SET_COOKIE, v);
}
resp
}
async fn sso_callback(
State(state): State<AppState>,
Path(slug): Path<String>,
headers: HeaderMap,
Query(params): Query<CallbackParams>,
) -> Response {
let Some(secret) = state.config.session_secret.as_ref() else {
return login_error(&state, "disabled");
};
if params.error.is_some() {
return login_error(&state, "denied");
}
let (Some(code), Some(cb_state)) = (params.code, params.state) else {
return login_error(&state, "callback");
};
let Some(sealed) = read_cookie(&headers, SSO_FLOW_COOKIE) else {
return login_error(&state, "expired");
};
let flow = match open_flow(&sealed, secret.key()) {
Ok(f) => f,
Err(_) => return login_error(&state, "expired"),
};
let Some(redirect_uri) = derive_bare_redirect(&headers, &state, &slug) else {
return login_error(&state, "config");
};
let cfg = match super::sso_provider::resolve_by_slug(&state.pool, &slug, redirect_uri).await {
Ok(Some(c)) => c,
Ok(None) => return login_error(&state, "disabled"),
Err(e) => {
tracing::error!(target: "rustango::admin::sso", "callback resolve: {e}");
return login_error(&state, "config");
}
};
let provider = match build_provider(&cfg).await {
Ok(p) => p,
Err(e) => {
tracing::error!(target: "rustango::admin::sso", "callback build: {e}");
return login_error(&state, "config");
}
};
let normalized = match provider.complete(&flow, &code, &cb_state).await {
Ok((u, _tokens)) => u,
Err(e) => {
tracing::warn!(target: "rustango::admin::sso", "handshake: {e}");
return login_error(&state, "handshake");
}
};
let email = match verified_email(&normalized) {
Ok(e) => e.to_ascii_lowercase(),
Err(_) => return login_error(&state, "unverified"),
};
let Some(user) = find_admin_user_by_email(&state.pool, &email).await else {
tracing::warn!(target: "rustango::admin::sso", "no admin account for {email}");
return login_error(&state, "nouser");
};
if !user.active {
return login_error(&state, "inactive");
}
let auth_hash = session::password_fingerprint(secret, &user.password_hash);
let cookie_value = session::encode(
secret,
AdminSession {
user_id: user.id,
username: user.username,
is_superuser: user.is_superuser,
},
&auth_hash,
);
let session_cookie = format!(
"{SESSION_COOKIE}={cookie_value}; Path=/; HttpOnly; SameSite=Lax{s}",
s = cookie_attrs(state.config.secure_cookies),
);
let clear_flow = format!("{SSO_FLOW_COOKIE}=; Path=/; HttpOnly; Max-Age=0");
let redirect_to = if state.config.admin_prefix.is_empty() {
"/".to_owned()
} else {
state.config.admin_prefix.clone()
};
let mut resp = Redirect::to(&redirect_to).into_response();
if let Ok(v) = HeaderValue::from_str(&session_cookie) {
resp.headers_mut().append(header::SET_COOKIE, v);
}
if let Ok(v) = HeaderValue::from_str(&clear_flow) {
resp.headers_mut().append(header::SET_COOKIE, v);
}
resp
}
struct LinkedAdmin {
id: i64,
username: String,
password_hash: String,
is_superuser: bool,
active: bool,
}
async fn find_admin_user_by_email(pool: &crate::sql::Pool, email: &str) -> Option<LinkedAdmin> {
use crate::core::{SelectQuery, SqlValue};
let select = SelectQuery::by_pk(
AdminUser::SCHEMA,
"email",
SqlValue::String(email.to_owned()),
);
let fields: Vec<&'static crate::core::FieldSchema> = AdminUser::SCHEMA.fields.iter().collect();
let row = crate::sql::select_one_row_as_json(pool, &select, &fields)
.await
.ok()
.flatten()?;
Some(LinkedAdmin {
id: row.get("id").and_then(serde_json::Value::as_i64)?,
username: row.get("username").and_then(|v| v.as_str())?.to_owned(),
password_hash: row
.get("password_hash")
.and_then(|v| v.as_str())?
.to_owned(),
is_superuser: row
.get("is_superuser")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false),
active: row
.get("active")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false),
})
}
#[cfg(test)]
mod tests {
use super::*;
fn cfg(provider: &str) -> ResolvedSso {
ResolvedSso {
provider: provider.into(),
issuer_url: None,
client_id: "cid".into(),
client_secret: "csecret".into(),
redirect_uri: "https://app.example.com/login/sso/x/callback".into(),
scopes: None,
}
}
#[tokio::test]
async fn presets_build_without_network() {
for name in ["google", "microsoft", "github", "gitlab", "discord"] {
let p = build_provider(&cfg(name)).await.expect("preset builds");
assert_eq!(p.client_id, "cid");
assert_eq!(
p.redirect_uri,
"https://app.example.com/login/sso/x/callback"
);
}
}
#[tokio::test]
async fn scopes_override_is_applied() {
let p = build_provider(&cfg("google")).await.expect("builds");
assert_eq!(p.scopes, vec!["openid", "email", "profile"]);
let mut c = cfg("google");
c.scopes = Some(vec!["openid".into(), "email".into(), "groups".into()]);
let p = build_provider(&c).await.expect("builds");
assert_eq!(p.scopes, vec!["openid", "email", "groups"]);
}
#[test]
fn parse_scopes_splits_and_trims() {
assert_eq!(parse_scopes(None), None);
assert_eq!(parse_scopes(Some(" ")), None);
assert_eq!(
parse_scopes(Some("openid email profile")),
Some(vec!["openid".into(), "email".into(), "profile".into()])
);
}
#[test]
fn resolve_secret_ref_env_reads_env_or_literal() {
assert_eq!(
resolve_secret_ref_env("plain-literal").unwrap(),
"plain-literal"
);
assert!(resolve_secret_ref_env("env://PATH").is_ok());
assert!(resolve_secret_ref_env("env://RUSTANGO_TEST_UNSET_VAR_QQQ").is_err());
}
#[tokio::test]
async fn unknown_provider_is_rejected() {
let e = build_provider(&cfg("myspace")).await.unwrap_err();
assert!(matches!(e, SsoError::UnknownProvider(p) if p == "myspace"));
}
#[tokio::test]
async fn oidc_without_issuer_is_rejected() {
assert!(matches!(
build_provider(&cfg("oidc")).await.unwrap_err(),
SsoError::MissingIssuer
));
}
#[tokio::test]
async fn empty_client_id_is_rejected() {
let mut c = cfg("google");
c.client_id = " ".into();
assert!(matches!(
build_provider(&c).await.unwrap_err(),
SsoError::Config(_)
));
}
#[test]
fn verified_email_requires_verified_flag() {
let mut u = NormalizedUser {
provider: "google".into(),
provider_user_id: "1".into(),
email: Some("a@example.com".into()),
email_verified: false,
name: None,
avatar_url: None,
raw: serde_json::json!({}),
};
assert!(verified_email(&u).is_err());
u.email_verified = true;
assert_eq!(verified_email(&u).unwrap(), "a@example.com");
}
}