use std::sync::Arc;
use axum::{
Json, Router,
extract::{Form, Query, State},
http::StatusCode,
response::{IntoResponse, Redirect, Response},
routing::{get, post},
};
use serde::Deserialize;
use super::{
SamlIdpConfig, SamlReplayCache, effective_saml_email_verified, registry::SamlIdpRegistry,
replay::SamlReplayStore, verify::verify_saml_response,
};
use crate::{
account_linking::AccountStore,
audit::logger::{AuditEventType, SecretType, get_audit_logger},
handlers::generate_secure_state,
session::{SessionStore, unix_now},
state_store::StateStore,
};
const RELAY_PAYLOAD_SEPARATOR: char = '\n';
const LOGIN_STATE_TTL_SECS: u64 = 600;
const SESSION_TTL_SECS: u64 = 7 * 24 * 60 * 60;
#[derive(Clone)]
pub struct SamlAuthState {
registry: SamlIdpRegistry,
state_store: Arc<dyn StateStore>,
session_store: Arc<dyn SessionStore>,
user_store: Option<Arc<dyn AccountStore>>,
replay: Arc<dyn SamlReplayStore>,
}
impl SamlAuthState {
#[must_use]
pub fn new(state_store: Arc<dyn StateStore>, session_store: Arc<dyn SessionStore>) -> Self {
Self {
registry: SamlIdpRegistry::new(),
state_store,
session_store,
user_store: None,
replay: Arc::new(SamlReplayCache::new()),
}
}
#[must_use]
pub fn with_replay_store(mut self, replay: Arc<dyn SamlReplayStore>) -> Self {
self.replay = replay;
self
}
#[must_use]
pub fn replay_is_distributed(&self) -> bool {
self.replay.is_distributed()
}
#[must_use]
pub fn with_idp(mut self, idp: SamlIdpConfig) -> Self {
self.registry = self.registry.with_config_idp(idp);
self
}
#[must_use]
pub fn with_registry(mut self, registry: SamlIdpRegistry) -> Self {
self.registry = registry;
self
}
#[must_use]
pub const fn registry(&self) -> &SamlIdpRegistry {
&self.registry
}
#[must_use]
pub fn with_user_store(mut self, user_store: Arc<dyn AccountStore>) -> Self {
self.user_store = Some(user_store);
self
}
#[must_use]
pub fn idp_names(&self) -> Vec<String> {
self.registry.idp_names()
}
}
pub fn saml_routes(state: SamlAuthState) -> Router {
Router::new()
.route("/auth/saml/login", get(saml_login))
.route("/auth/saml/acs", post(saml_acs))
.route("/auth/saml/metadata", get(saml_metadata))
.with_state(state)
}
#[derive(Debug, Deserialize)]
pub struct LoginQuery {
pub idp: String,
#[serde(default)]
pub tenant: Option<String>,
}
#[derive(Debug, Deserialize)]
pub struct AcsForm {
#[serde(rename = "SAMLResponse")]
pub saml_response: String,
#[serde(rename = "RelayState", default)]
pub relay_state: String,
}
fn json_error(status: StatusCode, message: &str) -> Response {
(status, Json(serde_json::json!({ "error": message }))).into_response()
}
#[derive(Debug, PartialEq, Eq)]
struct RelayPayload {
idp_name: String,
tenant: Option<String>,
request_id: String,
}
impl RelayPayload {
fn encode(&self) -> String {
format!(
"{}{RELAY_PAYLOAD_SEPARATOR}{}{RELAY_PAYLOAD_SEPARATOR}{}",
self.idp_name,
self.tenant.as_deref().unwrap_or_default(),
self.request_id
)
}
fn decode(payload: &str) -> Option<Self> {
let mut parts = payload.splitn(3, RELAY_PAYLOAD_SEPARATOR);
let idp_name = parts.next()?.to_string();
let tenant = parts.next()?;
let request_id = parts.next()?.to_string();
Some(Self {
idp_name,
tenant: (!tenant.is_empty()).then(|| tenant.to_string()),
request_id,
})
}
}
pub async fn saml_login(
State(state): State<SamlAuthState>,
Query(q): Query<LoginQuery>,
) -> Response {
let Some(idp) = state.registry.resolve(&q.idp, q.tenant.as_deref()) else {
return json_error(StatusCode::NOT_FOUND, "unknown SAML IdP");
};
let Some(sso_url) = idp.sso_redirect_url() else {
tracing::error!(idp = %q.idp, "IdP metadata has no HTTP-Redirect SSO endpoint");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "IdP configuration error");
};
let authn_request = match idp.service_provider().make_authentication_request(&sso_url) {
Ok(req) => req,
Err(e) => {
tracing::error!(error = %e, "failed to build AuthnRequest");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "could not start SAML login");
},
};
let relay_state = generate_secure_state();
let payload = RelayPayload {
idp_name: idp.idp_name.clone(),
tenant: idp.tenant_id.clone(),
request_id: authn_request.id.clone(),
}
.encode();
let Ok(now) = unix_now() else {
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "system clock error");
};
if let Err(e) = state
.state_store
.store(relay_state.clone(), payload, now + LOGIN_STATE_TTL_SECS)
.await
{
tracing::error!(error = %e, "failed to store SAML RelayState");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "could not start SAML login");
}
let redirect = match idp.signing_key() {
Some(key) => authn_request.signed_redirect(&relay_state, key),
None => authn_request.redirect(&relay_state),
};
match redirect {
Ok(Some(url)) => Redirect::to(url.as_str()).into_response(),
Ok(None) => {
json_error(StatusCode::INTERNAL_SERVER_ERROR, "IdP has no redirect destination")
},
Err(e) => {
tracing::error!(error = %e, "failed to build SAML redirect");
json_error(StatusCode::INTERNAL_SERVER_ERROR, "could not start SAML login")
},
}
}
pub async fn saml_metadata(
State(state): State<SamlAuthState>,
Query(q): Query<LoginQuery>,
) -> Response {
let Some(idp) = state.registry.resolve(&q.idp, q.tenant.as_deref()) else {
return json_error(StatusCode::NOT_FOUND, "unknown SAML IdP");
};
(
[(axum::http::header::CONTENT_TYPE, "application/samlmetadata+xml")],
idp.sp_metadata_xml(),
)
.into_response()
}
pub async fn saml_acs(State(state): State<SamlAuthState>, Form(form): Form<AcsForm>) -> Response {
let logger = get_audit_logger();
if form.relay_state.is_empty() {
return json_error(StatusCode::BAD_REQUEST, "missing RelayState");
}
let Ok((payload, expiry)) = state.state_store.retrieve(&form.relay_state).await else {
return json_error(StatusCode::BAD_REQUEST, "invalid or expired RelayState");
};
let Some(RelayPayload {
idp_name,
tenant,
request_id,
}) = RelayPayload::decode(&payload)
else {
return json_error(StatusCode::BAD_REQUEST, "malformed RelayState");
};
let Ok(now_secs) = unix_now() else {
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "system clock error");
};
if now_secs > expiry {
return json_error(StatusCode::BAD_REQUEST, "RelayState expired");
}
let Some(idp) = state.registry.resolve(&idp_name, tenant.as_deref()) else {
tracing::error!(
idp = %idp_name,
"RelayState referenced an IdP that no longer resolves for its tenant"
);
return json_error(StatusCode::BAD_REQUEST, "SAML authentication failed");
};
let idp = idp.as_ref();
let assertion = match verify_saml_response(
idp,
&form.saml_response,
&[request_id.as_str()],
state.replay.as_ref(),
chrono::Utc::now(),
)
.await
{
Ok(a) => a,
Err(e) => {
tracing::warn!(idp = %idp_name, error = %e, "SAML assertion verification failed");
logger.log_failure(
AuditEventType::AuthFailure,
SecretType::SessionToken,
None,
"saml_acs",
&format!("verification_failed:{idp_name}"),
);
let status = match e {
super::SamlError::Replay => StatusCode::UNAUTHORIZED,
_ => StatusCode::BAD_REQUEST,
};
return json_error(status, "SAML authentication failed");
},
};
let provider = idp.provider_key();
let email_verified = effective_saml_email_verified(idp);
let local_user_id = if let Some(store) = &state.user_store {
match store
.link_or_create_user(
assertion.email.as_deref(),
email_verified,
&provider,
&assertion.name_id,
)
.await
{
Ok(result) => result.user_id,
Err(e) => {
tracing::error!(error = %e, "account store lookup failed");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "user resolution failed");
},
}
} else {
format!("{provider}:{}", assertion.name_id)
};
let session = match state
.session_store
.create_session(&local_user_id, now_secs + SESSION_TTL_SECS)
.await
{
Ok(tokens) => tokens,
Err(e) => {
tracing::error!(error = %e, "session creation failed");
return json_error(StatusCode::INTERNAL_SERVER_ERROR, "session could not be created");
},
};
logger.log_success(
AuditEventType::AuthSuccess,
SecretType::SessionToken,
Some(local_user_id),
&format!("saml_acs:{idp_name}"),
);
Json(serde_json::json!({
"access_token": session.access_token,
"refresh_token": session.refresh_token,
"token_type": "Bearer",
"expires_in": session.expires_in,
"provider": provider,
}))
.into_response()
}