pub mod provider;
pub use provider::{list_enabled, resolve_by_slug, SsoProvider};
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),
}
}
#[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");
}
}