use crate::constants::{status_code, ParserType};
use crate::error::OpenSamlError;
use crate::xml::{extract, fields};
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 {
let now = OffsetDateTime::now_utc();
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<(), OpenSamlError> {
let fields = match parser_type {
ParserType::SamlResponse => fields::login_response_status_fields(),
ParserType::LogoutResponse => fields::logout_response_status_fields(),
_ => return Ok(()),
};
let result = extract(content, &fields)?;
match result.get_str("top") {
Some(code) if code == status_code::SUCCESS => Ok(()),
Some(code) if !code.is_empty() => Err(OpenSamlError::FailedStatus {
top: code.to_string(),
second: result.get_str("second").unwrap_or_default().to_string(),
}),
_ => Err(OpenSamlError::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(OpenSamlError::FailedStatus { top, second }) => {
assert_eq!(top, status_code::REQUESTER);
assert_eq!(second, status_code::INVALID_NAME_ID_POLICY);
}
other => return Err(format!("expected FailedStatus, got {other:?}").into()),
}
Ok(())
}
}