use std::collections::HashMap;
use chrono::{DateTime, Utc};
use super::{SamlError, SamlIdpConfig, replay::SamlReplayCache};
const NAMEID_FORMAT_EMAIL: &str = "urn:oasis:names:tc:SAML:1.1:nameid-format:emailAddress";
const FALLBACK_REPLAY_WINDOW_SECS: i64 = 300;
#[derive(Debug, Clone)]
pub struct VerifiedAssertion {
pub name_id: String,
pub name_id_format: Option<String>,
pub email: Option<String>,
pub display_name: Option<String>,
pub attributes: HashMap<String, Vec<String>>,
pub not_on_or_after: Option<DateTime<Utc>>,
}
pub fn reject_doctype(xml: &str) -> Result<(), SamlError> {
let lowered = xml.to_ascii_lowercase();
if lowered.contains("<!doctype") || lowered.contains("<!entity") {
return Err(SamlError::DocTypeForbidden);
}
Ok(())
}
pub fn verify_saml_response(
idp: &SamlIdpConfig,
response_b64: &str,
possible_request_ids: &[&str],
replay: &SamlReplayCache,
now: DateTime<Utc>,
) -> Result<VerifiedAssertion, SamlError> {
use base64::Engine as _;
let raw = base64::engine::general_purpose::STANDARD
.decode(response_b64.trim())
.map_err(|e| SamlError::Malformed(format!("base64 decode failed: {e}")))?;
let xml = std::str::from_utf8(&raw)
.map_err(|e| SamlError::Malformed(format!("response is not valid UTF-8: {e}")))?;
reject_doctype(xml)?;
let assertion = idp
.service_provider()
.parse_xml_response(xml, Some(possible_request_ids))
.map_err(|e| SamlError::Verification(e.to_string()))?;
if assertion.id.trim().is_empty() {
return Err(SamlError::MissingField("assertion ID"));
}
let not_on_or_after = assertion.conditions.as_ref().and_then(|c| c.not_on_or_after);
let replay_expiry = not_on_or_after
.unwrap_or_else(|| now + chrono::Duration::seconds(FALLBACK_REPLAY_WINDOW_SECS));
if !replay.check_and_record(&assertion.id, replay_expiry, now) {
return Err(SamlError::Replay);
}
let name_id_subject = assertion
.subject
.as_ref()
.and_then(|s| s.name_id.as_ref())
.ok_or(SamlError::MissingField("subject NameID"))?;
let name_id = name_id_subject.value.trim().to_string();
if name_id.is_empty() {
return Err(SamlError::MissingField("subject NameID"));
}
let name_id_format = name_id_subject.format.clone();
let mut attributes: HashMap<String, Vec<String>> = HashMap::new();
if let Some(statements) = &assertion.attribute_statements {
for statement in statements {
for attribute in &statement.attributes {
let Some(name) = attribute.name.clone() else {
continue;
};
let values = attribute.values.iter().filter_map(|v| v.value.clone());
attributes.entry(name).or_default().extend(values);
}
}
}
let email = first_nonempty(&attributes, &idp.attribute_mapping.email)
.map(str::to_string)
.or_else(|| {
(name_id_format.as_deref() == Some(NAMEID_FORMAT_EMAIL)).then(|| name_id.clone())
});
let display_name =
first_nonempty(&attributes, &idp.attribute_mapping.display_name).map(str::to_string);
Ok(VerifiedAssertion {
name_id,
name_id_format,
email,
display_name,
attributes,
not_on_or_after,
})
}
fn first_nonempty<'a>(
attributes: &'a HashMap<String, Vec<String>>,
names: &[String],
) -> Option<&'a str> {
names.iter().find_map(|name| {
attributes
.get(name)
.and_then(|values| values.iter().map(String::as_str).find(|v| !v.trim().is_empty()))
})
}