use uppsala::{Document, NodeId};
use crate::xml::error::XmlError;
pub trait SamlDeserialize<'a>: Sized {
fn from_xml(doc: &'a Document<'a>, node: NodeId) -> Result<Self, XmlError>;
}
pub fn parse_saml<'a, T: SamlDeserialize<'a>>(doc: &'a Document<'a>) -> Result<T, XmlError> {
let root = doc.document_element().ok_or(XmlError::EmptyDocument)?;
T::from_xml(doc, root)
}
pub fn parse_secure(xml: &str) -> Result<Document<'_>, uppsala::XmlError> {
let doc = uppsala::parse(xml)?;
if doc.doctype.is_some() {
let (line, column) = locate_doctype(xml).unwrap_or((1, 1));
return Err(uppsala::XmlError::well_formedness(
"DOCTYPE/DTD declarations are forbidden in SAML messages",
line,
column,
));
}
Ok(doc)
}
fn locate_doctype(xml: &str) -> Option<(usize, usize)> {
let offset = xml.find("<!DOCTYPE")?;
let mut line = 1usize;
let mut line_start = 0usize;
for (i, b) in xml.as_bytes()[..offset].iter().enumerate() {
if *b == b'\n' {
line += 1;
line_start = i + 1;
}
}
let column = xml[line_start..offset].chars().count() + 1;
Some((line, column))
}
#[cfg(test)]
mod parse_secure_tests {
use super::{locate_doctype, parse_secure};
#[test]
fn reports_doctype_position() {
assert_eq!(
locate_doctype("<?xml version=\"1.0\"?>\n<!DOCTYPE x [ ]>\n<x/>"),
Some((2, 1))
);
assert_eq!(locate_doctype(" <!DOCTYPE x><x/>"), Some((1, 4)));
assert_eq!(
locate_doctype("<!-- café -->\n<!DOCTYPE x><x/>"),
Some((2, 1))
);
assert_eq!(locate_doctype("<x/>"), None);
}
#[test]
fn rejects_doctype_declaration() {
let xml = r#"<?xml version="1.0"?>
<!DOCTYPE samlp:Response [ <!ENTITY x "expanded"> ]>
<samlp:Response xmlns:samlp="urn:oasis:names:tc:SAML:2.0:protocol">&x;</samlp:Response>"#;
assert!(
uppsala::parse(xml).is_ok(),
"precondition: the DTD-bearing document is itself well-formed"
);
assert!(
parse_secure(xml).is_err(),
"parse_secure must reject the document solely because of the DTD"
);
}
#[test]
fn rejects_internal_subset_without_entities() {
let xml = r#"<!DOCTYPE Response><Response/>"#;
assert!(parse_secure(xml).is_err());
}
#[test]
fn accepts_well_formed_saml_without_dtd() {
let xml = r#"<samlp:Response xmlns:samlp="urn:oasis:names:tc:SAML:2.0:protocol" ID="_1"/>"#;
let doc = parse_secure(xml).expect("DTD-free SAML must parse");
assert!(doc.document_element().is_some());
}
}