use std::{collections::HashMap, sync::Arc, time::Duration};
use arc_swap::ArcSwap;
use chrono::{DateTime, Utc};
use super::{
SamlError, SamlIdpConfig,
store::{SamlIdpRecord, SamlIdpSpec, SamlIdpStore},
};
pub const DEFAULT_REFRESH_INTERVAL: Duration = Duration::from_secs(30);
pub const DEFAULT_EXPIRY_WARNING_DAYS: i64 = 30;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum IdpSource {
ConfigFile,
Store,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CertExpiryWarning {
pub idp_name: String,
pub tenant_id: Option<String>,
pub expires_at: DateTime<Utc>,
pub expired: bool,
}
#[derive(Clone, Default)]
pub struct SpKeyMaterial {
pub private_key: Vec<u8>,
pub certificate: Vec<u8>,
pub sign_authn_requests: bool,
pub previous_private_key: Option<Vec<u8>>,
pub previous_certificate: Option<Vec<u8>>,
}
impl std::fmt::Debug for SpKeyMaterial {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SpKeyMaterial")
.field("sign_authn_requests", &self.sign_authn_requests)
.field("has_previous_key", &self.previous_private_key.is_some())
.finish_non_exhaustive()
}
}
#[derive(Clone)]
pub struct SamlIdpRegistry {
config_idps: Arc<HashMap<String, Arc<SamlIdpConfig>>>,
store: Option<Arc<dyn SamlIdpStore>>,
cached: Arc<ArcSwap<HashMap<String, Arc<SamlIdpConfig>>>>,
sp_keys: Option<Arc<SpKeyMaterial>>,
}
impl std::fmt::Debug for SamlIdpRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SamlIdpRegistry")
.field("config_idps", &self.config_idps.keys().collect::<Vec<_>>())
.field("has_store", &self.store.is_some())
.field("stored", &self.cached.load().keys().cloned().collect::<Vec<_>>())
.field("sp_keys", &self.sp_keys)
.finish()
}
}
impl Default for SamlIdpRegistry {
fn default() -> Self {
Self::new()
}
}
fn normalize_tenant(tenant: Option<&str>) -> Option<&str> {
tenant.map(str::trim).filter(|t| !t.is_empty())
}
impl SamlIdpRegistry {
#[must_use]
pub fn new() -> Self {
Self {
config_idps: Arc::new(HashMap::new()),
store: None,
cached: Arc::new(ArcSwap::from_pointee(HashMap::new())),
sp_keys: None,
}
}
#[must_use]
pub fn with_sp_keys(mut self, sp_keys: SpKeyMaterial) -> Self {
self.sp_keys = Some(Arc::new(sp_keys));
self
}
#[must_use]
pub fn with_config_idp(mut self, idp: SamlIdpConfig) -> Self {
let idps = Arc::make_mut(&mut self.config_idps);
idps.insert(idp.idp_name.clone(), Arc::new(idp));
self
}
#[must_use]
pub fn with_store(mut self, store: Arc<dyn SamlIdpStore>) -> Self {
self.store = Some(store);
self
}
#[must_use]
pub fn has_store(&self) -> bool {
self.store.is_some()
}
#[must_use]
pub fn resolve(&self, idp_name: &str, tenant: Option<&str>) -> Option<Arc<SamlIdpConfig>> {
let idp = self
.config_idps
.get(idp_name)
.cloned()
.or_else(|| self.cached.load().get(idp_name).cloned())?;
(normalize_tenant(idp.tenant_id.as_deref()) == normalize_tenant(tenant)).then_some(idp)
}
#[must_use]
pub fn idp_names(&self) -> Vec<String> {
let mut names: Vec<String> = self
.config_idps
.keys()
.cloned()
.chain(self.cached.load().keys().cloned())
.collect();
names.sort_unstable();
names.dedup();
names
}
#[must_use]
pub fn source_of(&self, idp_name: &str) -> Option<IdpSource> {
if self.config_idps.contains_key(idp_name) {
return Some(IdpSource::ConfigFile);
}
self.cached.load().contains_key(idp_name).then_some(IdpSource::Store)
}
pub async fn refresh(&self) -> Result<(), SamlError> {
let Some(store) = &self.store else {
return Ok(());
};
let records = store.list().await?;
let mut next: HashMap<String, Arc<SamlIdpConfig>> = HashMap::with_capacity(records.len());
for record in records {
if self.config_idps.contains_key(&record.idp_name) {
tracing::error!(
idp = %record.idp_name,
"stored SAML IdP shadows a [saml.idps.*] entry of the same name and will \
NOT be served: both would share the saml:<name> account namespace. \
Rename the stored IdP or remove the config-file entry."
);
continue;
}
match config_from_record(&record, self.sp_keys.as_deref()) {
Ok(config) => {
next.insert(record.idp_name.clone(), Arc::new(config));
},
Err(e) => {
tracing::error!(
idp = %record.idp_name, error = %e,
"stored SAML IdP failed to build and will NOT be served"
);
},
}
}
self.cached.store(Arc::new(next));
Ok(())
}
pub async fn create(&self, spec: &SamlIdpSpec) -> Result<SamlIdpRecord, SamlError> {
let store = self.require_store()?;
if self.config_idps.contains_key(&spec.idp_name) {
return Err(SamlError::NameTaken(spec.idp_name.clone()));
}
let record = store.create(spec).await?;
self.refresh().await?;
Ok(record)
}
pub async fn update(&self, spec: &SamlIdpSpec) -> Result<SamlIdpRecord, SamlError> {
let record = self.require_store()?.update(spec).await?;
self.refresh().await?;
Ok(record)
}
pub async fn delete(&self, idp_name: &str) -> Result<(), SamlError> {
self.require_store()?.delete(idp_name).await?;
self.refresh().await
}
pub async fn list_stored(&self) -> Result<Vec<SamlIdpRecord>, SamlError> {
self.require_store()?.list().await
}
pub async fn get_stored(&self, idp_name: &str) -> Result<Option<SamlIdpRecord>, SamlError> {
self.require_store()?.get(idp_name).await
}
fn require_store(&self) -> Result<&Arc<dyn SamlIdpStore>, SamlError> {
self.store.as_ref().ok_or_else(|| {
SamlError::Store(
"no SAML IdP store is configured — set [saml] store_enabled = true on a \
deployment with a database pool"
.to_string(),
)
})
}
#[must_use]
pub fn expiry_report(&self, now: DateTime<Utc>, threshold_days: i64) -> Vec<CertExpiryWarning> {
let horizon = now + chrono::Duration::days(threshold_days);
let stored = self.cached.load();
let mut warnings: Vec<CertExpiryWarning> = self
.config_idps
.values()
.chain(stored.values())
.filter_map(|idp| {
let expires_at = idp.signing_certificate_expiry()?;
(expires_at <= horizon).then(|| CertExpiryWarning {
idp_name: idp.idp_name.clone(),
tenant_id: idp.tenant_id.clone(),
expires_at,
expired: expires_at <= now,
})
})
.collect();
warnings.sort_by(|a, b| a.expires_at.cmp(&b.expires_at).then(a.idp_name.cmp(&b.idp_name)));
warnings
}
pub fn log_expiry_report(&self, threshold_days: i64) {
for warning in self.expiry_report(Utc::now(), threshold_days) {
if warning.expired {
tracing::error!(
idp = %warning.idp_name, expired_at = %warning.expires_at,
"SAML IdP signing certificate has EXPIRED — SSO for this IdP is down \
until the operator loads fresh metadata"
);
} else {
tracing::warn!(
idp = %warning.idp_name, expires_at = %warning.expires_at,
"SAML IdP signing certificate expires soon — load fresh metadata before \
the cliff"
);
}
}
}
pub async fn refresh_loop(self, interval: Duration, warning_days: i64) {
let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
ticker.tick().await;
if let Err(e) = self.refresh().await {
tracing::error!(
error = %e,
"SAML IdP refresh failed; continuing to serve the last good generation"
);
}
self.log_expiry_report(warning_days);
}
}
}
fn config_from_record(
record: &SamlIdpRecord,
sp_keys: Option<&SpKeyMaterial>,
) -> Result<SamlIdpConfig, SamlError> {
let mut builder = SamlIdpConfig::builder(
record.idp_name.clone(),
record.sp_entity_id.clone(),
record.acs_url.clone(),
)
.idp_metadata_xml(&record.metadata_xml)?
.tenant_id(record.tenant_id.map(|t| t.to_string()))
.trust_asserted_email(record.trust_asserted_email);
if let Some(keys) = sp_keys {
builder = builder
.sp_key_pair(&keys.private_key, &keys.certificate)?
.sign_authn_requests(keys.sign_authn_requests);
if let (Some(key), Some(cert)) =
(keys.previous_private_key.as_deref(), keys.previous_certificate.as_deref())
{
builder = builder.sp_previous_key_pair(key, cert)?;
}
}
builder.build()
}