use crate::constants::{status_code, ParserType};
use crate::error::SamlError;
use crate::xml::{extract_with_limits, fields, XmlLimits};
use time::{format_description::well_known::Rfc3339, Duration, OffsetDateTime};
fn parse(ts: &str) -> Option<OffsetDateTime> {
OffsetDateTime::parse(ts, &Rfc3339).ok()
}
pub fn verify_time(
not_before: Option<&str>,
not_on_or_after: Option<&str>,
drift: (i64, i64),
) -> bool {
verify_time_at(
not_before,
not_on_or_after,
drift,
OffsetDateTime::now_utc(),
)
}
pub(crate) fn verify_time_at(
not_before: Option<&str>,
not_on_or_after: Option<&str>,
drift: (i64, i64),
now: OffsetDateTime,
) -> bool {
let (nb_drift, na_drift) = (
Duration::milliseconds(drift.0),
Duration::milliseconds(drift.1),
);
match (not_before, not_on_or_after) {
(None, None) => true,
(Some(nb), None) => match parse(nb) {
Some(t) => t + nb_drift <= now,
None => false,
},
(None, Some(na)) => match parse(na) {
Some(t) => now < t + na_drift,
None => false,
},
(Some(nb), Some(na)) => match (parse(nb), parse(na)) {
(Some(b), Some(a)) => b + nb_drift <= now && now < a + na_drift,
_ => false,
},
}
}
pub fn check_status(content: &str, parser_type: ParserType) -> Result<(), SamlError> {
check_status_with_limits(content, parser_type, XmlLimits::default())
}
pub fn check_status_with_limits(
content: &str,
parser_type: ParserType,
limits: XmlLimits,
) -> Result<(), SamlError> {
let fields = match parser_type {
ParserType::SamlResponse => fields::login_response_status_fields(),
ParserType::LogoutResponse => fields::logout_response_status_fields(),
_ => return Ok(()),
};
let result = extract_with_limits(content, &fields, limits)?;
match result.get_str("top") {
Some(code) if code == status_code::SUCCESS => Ok(()),
Some(code) if !code.is_empty() => Err(SamlError::StatusNotSuccess {
top: code.to_string(),
second: result
.get_str("second")
.filter(|second| !second.is_empty())
.map(str::to_string),
}),
_ => Err(SamlError::UndefinedStatus),
}
}
#[cfg(test)]
mod tests {
use super::*;
const RESPONSE: &str = include_str!("../tests/fixtures/response.xml");
const FAILED: &str = include_str!("../tests/fixtures/failed_response.xml");
#[test]
fn time_window_basic() {
assert!(verify_time(None, None, (0, 0)));
assert!(verify_time(
Some("2000-01-01T00:00:00Z"),
Some("2999-01-01T00:00:00Z"),
(0, 0)
));
assert!(!verify_time(None, Some("2000-01-01T00:00:00Z"), (0, 0)));
assert!(!verify_time(Some("2999-01-01T00:00:00Z"), None, (0, 0)));
assert!(!verify_time(Some("not-a-date"), None, (0, 0)));
}
#[test]
fn drift_widens_window() {
assert!(verify_time(
None,
Some("2000-01-01T00:00:00Z"),
(0, 9_000_000_000_000)
));
assert!(verify_time(
Some("2999-01-01T00:00:00Z"),
None,
(-50_000_000_000_000, 0)
));
}
#[test]
fn status_success_and_two_tier_failure() -> Result<(), Box<dyn std::error::Error>> {
check_status(RESPONSE, ParserType::SamlResponse)?;
check_status(RESPONSE, ParserType::SamlRequest)?;
match check_status(FAILED, ParserType::SamlResponse) {
Err(SamlError::StatusNotSuccess { top, second }) => {
assert_eq!(top, status_code::REQUESTER);
assert_eq!(second.as_deref(), Some(status_code::INVALID_NAME_ID_POLICY));
}
other => return Err(format!("expected StatusNotSuccess, got {other:?}").into()),
}
Ok(())
}
}