use base64::Engine as _;
use chrono::{DateTime, Utc};
use samael::{
crypto::AllowedSignatureAlgorithm,
metadata::EntityDescriptor,
service_provider::{ServiceProvider, ServiceProviderBuilder},
};
use super::SamlError;
const DEFAULT_EMAIL_ATTRS: &[&str] = &[
"urn:oid:0.9.2342.19200300.100.1.3",
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/emailaddress",
"email",
"mail",
"emailAddress",
];
const DEFAULT_NAME_ATTRS: &[&str] = &[
"urn:oid:2.16.840.1.113730.3.1.241",
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/name",
"displayName",
"name",
"cn",
];
fn default_allowed_algorithms() -> Vec<AllowedSignatureAlgorithm> {
vec![
AllowedSignatureAlgorithm::RsaSha256,
AllowedSignatureAlgorithm::RsaSha384,
AllowedSignatureAlgorithm::RsaSha512,
AllowedSignatureAlgorithm::EcdsaSha256,
AllowedSignatureAlgorithm::EcdsaSha384,
AllowedSignatureAlgorithm::EcdsaSha512,
]
}
#[derive(Debug, Clone)]
pub struct SamlAttributeMapping {
pub email: Vec<String>,
pub display_name: Vec<String>,
}
impl Default for SamlAttributeMapping {
fn default() -> Self {
Self {
email: DEFAULT_EMAIL_ATTRS.iter().map(|s| (*s).to_string()).collect(),
display_name: DEFAULT_NAME_ATTRS.iter().map(|s| (*s).to_string()).collect(),
}
}
}
const ALLOWED_KEY_TRANSPORT_ALGORITHMS: &[&str] = &[
"http://www.w3.org/2001/04/xmlenc#rsa-oaep-mgf1p",
"http://www.w3.org/2009/xmlenc11#rsa-oaep",
];
const ALLOWED_CONTENT_ENCRYPTION_ALGORITHMS: &[&str] = &[
"http://www.w3.org/2009/xmlenc11#aes128-gcm",
"http://www.w3.org/2009/xmlenc11#aes192-gcm",
"http://www.w3.org/2009/xmlenc11#aes256-gcm",
];
#[must_use]
pub fn key_transport_algorithm_allowed(algorithm: &str) -> bool {
ALLOWED_KEY_TRANSPORT_ALGORITHMS.contains(&algorithm)
}
#[must_use]
pub fn content_encryption_algorithm_allowed(algorithm: &str) -> bool {
ALLOWED_CONTENT_ENCRYPTION_ALGORITHMS.contains(&algorithm)
}
struct SpKeyPair {
key: openssl::pkey::PKey<openssl::pkey::Private>,
cert_der: Vec<u8>,
}
impl std::fmt::Debug for SpKeyPair {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SpKeyPair").finish_non_exhaustive()
}
}
fn parse_private_key(
bytes: &[u8],
) -> Result<openssl::pkey::PKey<openssl::pkey::Private>, SamlError> {
openssl::pkey::PKey::private_key_from_pem(bytes)
.or_else(|_| openssl::pkey::PKey::private_key_from_der(bytes))
.map_err(|e| SamlError::Config(format!("SP private key is neither valid PEM nor DER: {e}")))
}
fn parse_certificate_der(bytes: &[u8]) -> Result<Vec<u8>, SamlError> {
let cert = openssl::x509::X509::from_pem(bytes)
.or_else(|_| openssl::x509::X509::from_der(bytes))
.map_err(|e| {
SamlError::Config(format!("SP certificate is neither valid PEM nor DER: {e}"))
})?;
cert.to_der()
.map_err(|e| SamlError::Config(format!("SP certificate could not be re-encoded: {e}")))
}
pub struct SamlIdpConfig {
pub idp_name: String,
pub tenant_id: Option<String>,
pub trust_asserted_email: bool,
pub attribute_mapping: SamlAttributeMapping,
pub sign_authn_requests: bool,
pub(crate) sp: ServiceProvider,
sp_certificate_der: Option<Vec<u8>>,
sp_previous: Option<SpKeyPair>,
}
impl std::fmt::Debug for SamlIdpConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SamlIdpConfig")
.field("idp_name", &self.idp_name)
.field("tenant_id", &self.tenant_id)
.field("trust_asserted_email", &self.trust_asserted_email)
.field("attribute_mapping", &self.attribute_mapping)
.field("sign_authn_requests", &self.sign_authn_requests)
.field("has_sp_key", &self.sp.key.is_some())
.field("has_previous_sp_key", &self.sp_previous.is_some())
.finish_non_exhaustive()
}
}
impl SamlIdpConfig {
#[must_use]
pub fn builder(
idp_name: impl Into<String>,
sp_entity_id: impl Into<String>,
acs_url: impl Into<String>,
) -> SamlIdpConfigBuilder {
SamlIdpConfigBuilder {
idp_name: idp_name.into(),
sp_entity_id: sp_entity_id.into(),
acs_url: acs_url.into(),
idp_metadata: None,
tenant_id: None,
trust_asserted_email: false,
attribute_mapping: SamlAttributeMapping::default(),
sp_key_pair: None,
sp_previous_key_pair: None,
sign_authn_requests: false,
}
}
#[must_use]
pub fn provider_key(&self) -> String {
super::saml_provider_key(&self.idp_name)
}
#[must_use]
pub fn sso_redirect_url(&self) -> Option<String> {
self.sp.sso_binding_location(samael::metadata::HTTP_REDIRECT_BINDING)
}
#[must_use]
pub fn idp_entity_id(&self) -> &str {
self.sp.idp_metadata.entity_id.as_deref().unwrap_or_default()
}
#[must_use]
pub fn signing_certificate_expiry(&self) -> Option<DateTime<Utc>> {
let certs = self.sp.idp_signing_certs().ok().flatten()?;
certs.iter().filter_map(|cert| certificate_not_after(cert.der_data())).min()
}
#[must_use]
pub const fn has_sp_key(&self) -> bool {
self.sp.key.is_some()
}
#[must_use]
pub const fn signs_authn_requests(&self) -> bool {
self.sign_authn_requests && self.sp.key.is_some()
}
#[must_use]
pub fn sp_metadata_xml(&self) -> String {
use std::fmt::Write as _;
let entity_id = xml_escape(self.sp.entity_id.as_deref().unwrap_or_default());
let acs_url = xml_escape(self.sp.acs_url.as_deref().unwrap_or_default());
let signed = if self.signs_authn_requests() {
"true"
} else {
"false"
};
let mut descriptors = String::new();
if let Some(der) = self.sp_certificate_der.as_deref() {
let _ = write!(descriptors, "{}", key_descriptor("signing", der));
let _ = write!(descriptors, "{}", key_descriptor("encryption", der));
}
if let Some(previous) = &self.sp_previous {
let _ = write!(descriptors, "{}", key_descriptor("encryption", &previous.cert_der));
}
format!(
r#"<EntityDescriptor xmlns="urn:oasis:names:tc:SAML:2.0:metadata" entityID="{entity_id}">
<SPSSODescriptor protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol" AuthnRequestsSigned="{signed}" WantAssertionsSigned="true">
{descriptors} <AssertionConsumerService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST" Location="{acs_url}" index="0" isDefault="true"/>
</SPSSODescriptor>
</EntityDescriptor>"#
)
}
pub(crate) const fn signing_key(&self) -> Option<&openssl::pkey::PKey<openssl::pkey::Private>> {
if self.sign_authn_requests {
self.sp.key.as_ref()
} else {
None
}
}
pub(crate) fn service_provider_with_previous_key(&self) -> Option<ServiceProvider> {
let previous = self.sp_previous.as_ref()?;
let mut sp = self.sp.clone();
sp.key = Some(previous.key.clone());
Some(sp)
}
pub(crate) const fn service_provider(&self) -> &ServiceProvider {
&self.sp
}
}
fn key_descriptor(key_use: &str, cert_der: &[u8]) -> String {
let cert_b64 = base64::engine::general_purpose::STANDARD.encode(cert_der);
format!(
" <KeyDescriptor use=\"{key_use}\">\n \
<KeyInfo xmlns=\"http://www.w3.org/2000/09/xmldsig#\">\n \
<X509Data><X509Certificate>{cert_b64}</X509Certificate></X509Data>\n \
</KeyInfo>\n </KeyDescriptor>\n"
)
}
fn xml_escape(raw: &str) -> String {
raw.replace('&', "&")
.replace('<', "<")
.replace('>', ">")
.replace('"', """)
.replace('\'', "'")
}
fn certificate_not_after(der: &[u8]) -> Option<DateTime<Utc>> {
let cert = openssl::x509::X509::from_der(der).ok()?;
let now = openssl::asn1::Asn1Time::days_from_now(0).ok()?;
let diff = now.diff(cert.not_after()).ok()?;
Some(
Utc::now()
+ chrono::Duration::days(i64::from(diff.days))
+ chrono::Duration::seconds(i64::from(diff.secs)),
)
}
#[derive(Debug)]
pub struct SamlIdpConfigBuilder {
idp_name: String,
sp_entity_id: String,
acs_url: String,
idp_metadata: Option<EntityDescriptor>,
tenant_id: Option<String>,
trust_asserted_email: bool,
attribute_mapping: SamlAttributeMapping,
sp_key_pair: Option<SpKeyPair>,
sp_previous_key_pair: Option<SpKeyPair>,
sign_authn_requests: bool,
}
impl SamlIdpConfigBuilder {
pub fn idp_metadata_xml(mut self, xml: &str) -> Result<Self, SamlError> {
let descriptor: EntityDescriptor = xml
.parse()
.map_err(|e| SamlError::Config(format!("invalid IdP metadata XML: {e}")))?;
self.idp_metadata = Some(descriptor);
Ok(self)
}
pub fn idp_parts(
self,
idp_entity_id: &str,
sso_redirect_url: &str,
signing_cert_der: &[u8],
) -> Result<Self, SamlError> {
let xml = idp_metadata_xml_from_parts(idp_entity_id, sso_redirect_url, signing_cert_der);
self.idp_metadata_xml(&xml)
}
#[must_use]
pub fn tenant_id(mut self, tenant_id: Option<String>) -> Self {
self.tenant_id = tenant_id;
self
}
#[must_use]
pub const fn trust_asserted_email(mut self, trust: bool) -> Self {
self.trust_asserted_email = trust;
self
}
#[must_use]
pub fn attribute_mapping(mut self, mapping: SamlAttributeMapping) -> Self {
self.attribute_mapping = mapping;
self
}
pub fn sp_key_pair(
mut self,
private_key: &[u8],
certificate: &[u8],
) -> Result<Self, SamlError> {
self.sp_key_pair = Some(build_key_pair(private_key, certificate, "sp_key_pair")?);
Ok(self)
}
pub fn sp_previous_key_pair(
mut self,
private_key: &[u8],
certificate: &[u8],
) -> Result<Self, SamlError> {
self.sp_previous_key_pair =
Some(build_key_pair(private_key, certificate, "sp_previous_key_pair")?);
Ok(self)
}
#[must_use]
pub const fn sign_authn_requests(mut self, sign: bool) -> Self {
self.sign_authn_requests = sign;
self
}
pub fn build(self) -> Result<SamlIdpConfig, SamlError> {
let idp_metadata = self
.idp_metadata
.ok_or_else(|| SamlError::Config("IdP metadata not supplied".to_string()))?;
if self.sign_authn_requests && self.sp_key_pair.is_none() {
return Err(SamlError::Config(
"sign_authn_requests is on but no SP key pair is configured — an IdP that \
requires signed AuthnRequests would reject every login, and sending them \
unsigned instead would be a silent downgrade"
.to_string(),
));
}
let sp_certificate_der = self.sp_key_pair.as_ref().map(|kp| kp.cert_der.clone());
let mut builder = ServiceProviderBuilder::default();
builder
.entity_id(Some(self.sp_entity_id))
.acs_url(Some(self.acs_url))
.idp_metadata(idp_metadata)
.allowed_signature_algorithms(Some(default_allowed_algorithms()))
.allow_idp_initiated(false);
if let Some(key_pair) = self.sp_key_pair {
builder.key(Some(key_pair.key));
}
let sp = builder
.build()
.map_err(|e| SamlError::Config(format!("service provider build failed: {e}")))?;
let config = SamlIdpConfig {
idp_name: self.idp_name,
tenant_id: self.tenant_id,
trust_asserted_email: self.trust_asserted_email,
attribute_mapping: self.attribute_mapping,
sign_authn_requests: self.sign_authn_requests,
sp,
sp_certificate_der,
sp_previous: self.sp_previous_key_pair,
};
if config.sso_redirect_url().is_none() {
return Err(SamlError::Config(
"IdP metadata declares no HTTP-Redirect SingleSignOnService binding — \
SP-initiated login would be impossible. Check the metadata XML."
.to_string(),
));
}
Ok(config)
}
}
fn build_key_pair(
private_key: &[u8],
certificate: &[u8],
field: &str,
) -> Result<SpKeyPair, SamlError> {
let key = parse_private_key(private_key)?;
let cert_der = parse_certificate_der(certificate)?;
let cert = openssl::x509::X509::from_der(&cert_der)
.map_err(|e| SamlError::Config(format!("SP certificate could not be re-read: {e}")))?;
if !cert.public_key().is_ok_and(|public| public.public_eq(&key)) {
return Err(SamlError::Config(format!(
"{field}: the SP certificate does not match the SP private key"
)));
}
Ok(SpKeyPair { key, cert_der })
}
fn idp_metadata_xml_from_parts(
idp_entity_id: &str,
sso_redirect_url: &str,
signing_cert_der: &[u8],
) -> String {
let cert_b64 = base64::engine::general_purpose::STANDARD.encode(signing_cert_der);
format!(
r#"<EntityDescriptor xmlns="urn:oasis:names:tc:SAML:2.0:metadata" entityID="{idp_entity_id}">
<IDPSSODescriptor protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol">
<KeyDescriptor use="signing">
<KeyInfo xmlns="http://www.w3.org/2000/09/xmldsig#">
<X509Data><X509Certificate>{cert_b64}</X509Certificate></X509Data>
</KeyInfo>
</KeyDescriptor>
<SingleSignOnService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-Redirect" Location="{sso_redirect_url}"/>
</IDPSSODescriptor>
</EntityDescriptor>"#
)
}