use crate::bindings::error::BindingError;
use crate::bindings::soap::SOAP11_ACTOR_NEXT;
pub const PAOS_CONTENT_TYPE: &str = "application/vnd.paos+xml";
pub const PAOS_HEADER_VALUE: &str = "ver=\"urn:liberty:paos:2003-08\"";
pub const PAOS_NS: &str = "urn:liberty:paos:2003-08";
pub const ECP_NS: &str = "urn:oasis:names:tc:SAML:2.0:profiles:SSO:ecp";
#[derive(Debug, Clone)]
pub struct PaosRequest {
pub response_consumer_url: String,
pub service: Option<String>,
pub message_id: Option<String>,
}
#[derive(Debug, Clone)]
pub struct PaosResponse {
pub ref_to_message_id: Option<String>,
}
#[derive(Debug, Clone)]
pub struct EcpRequest {
pub issuer: Option<String>,
pub provider_name: Option<String>,
pub is_passive: bool,
pub idp_list: Vec<String>,
}
#[derive(Debug, Clone)]
pub struct EcpResponse {
pub assertion_consumer_service_url: String,
}
#[derive(Debug, Clone)]
pub struct EcpRelayState {
pub relay_state: String,
}
fn xml_escape(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for c in s.chars() {
match c {
'&' => out.push_str("&"),
'<' => out.push_str("<"),
'>' => out.push_str(">"),
'"' => out.push_str("""),
'\'' => out.push_str("'"),
_ => out.push(c),
}
}
out
}
pub fn paos_request_header_xml(req: &PaosRequest) -> String {
let mut xml = String::with_capacity(256);
xml.push_str(&format!(
r#"<paos:Request xmlns:paos="{}" soap:mustUnderstand="1" soap:actor="{}" responseConsumerURL="{}""#,
PAOS_NS,
SOAP11_ACTOR_NEXT,
xml_escape(&req.response_consumer_url)
));
if let Some(ref svc) = req.service {
xml.push_str(&format!(r#" service="{}""#, xml_escape(svc)));
}
if let Some(ref mid) = req.message_id {
xml.push_str(&format!(r#" messageID="{}""#, xml_escape(mid)));
}
xml.push_str("/>");
xml
}
pub fn paos_response_header_xml(resp: &PaosResponse) -> String {
let mut xml = String::with_capacity(128);
xml.push_str(&format!(
r#"<paos:Response xmlns:paos="{}" soap:mustUnderstand="1" soap:actor="{}""#,
PAOS_NS, SOAP11_ACTOR_NEXT
));
if let Some(ref mid) = resp.ref_to_message_id {
xml.push_str(&format!(r#" refToMessageID="{}""#, mid));
}
xml.push_str("/>");
xml
}
pub fn ecp_request_header_xml(req: &EcpRequest) -> String {
let mut xml = String::with_capacity(256);
xml.push_str(&format!(
r#"<ecp:Request xmlns:ecp="{}" soap:mustUnderstand="1" soap:actor="{}" IsPassive="{}""#,
ECP_NS, SOAP11_ACTOR_NEXT, req.is_passive
));
if let Some(ref pn) = req.provider_name {
xml.push_str(&format!(r#" ProviderName="{}""#, xml_escape(pn)));
}
if req.issuer.is_none() && req.idp_list.is_empty() {
xml.push_str("/>");
return xml;
}
xml.push('>');
if let Some(ref issuer) = req.issuer {
xml.push_str(&format!(
r#"<saml:Issuer xmlns:saml="urn:oasis:names:tc:SAML:2.0:assertion">{}</saml:Issuer>"#,
xml_escape(issuer)
));
}
if !req.idp_list.is_empty() {
xml.push_str(r#"<samlp:IDPList xmlns:samlp="urn:oasis:names:tc:SAML:2.0:protocol">"#);
for idp in &req.idp_list {
xml.push_str(&format!(
r#"<samlp:IDPEntry ProviderID="{}"/>"#,
xml_escape(idp)
));
}
xml.push_str("</samlp:IDPList>");
}
xml.push_str("</ecp:Request>");
xml
}
pub fn ecp_response_header_xml(resp: &EcpResponse) -> String {
format!(
r#"<ecp:Response xmlns:ecp="{}" soap:mustUnderstand="1" soap:actor="{}" AssertionConsumerServiceURL="{}"/>"#,
ECP_NS,
SOAP11_ACTOR_NEXT,
xml_escape(&resp.assertion_consumer_service_url)
)
}
pub fn ecp_relay_state_header_xml(rs: &EcpRelayState) -> String {
format!(
r#"<ecp:RelayState xmlns:ecp="{}" soap:mustUnderstand="1" soap:actor="{}">{}</ecp:RelayState>"#,
ECP_NS,
SOAP11_ACTOR_NEXT,
xml_escape(&rs.relay_state)
)
}
pub fn is_paos_request(request: &impl crate::bindings::traits::HttpRequest) -> bool {
let accept = request.header("Accept").unwrap_or("");
let paos = request.header("PAOS").unwrap_or("");
accept.contains(PAOS_CONTENT_TYPE) && paos.contains("urn:liberty:paos:2003-08")
}
pub fn build_ecp_phase1_envelope(
authn_request_xml: &str,
ecp_request: &EcpRequest,
paos_request: &PaosRequest,
relay_state: Option<&EcpRelayState>,
) -> Result<String, BindingError> {
let mut headers = String::new();
headers.push_str(&paos_request_header_xml(paos_request));
headers.push_str(&ecp_request_header_xml(ecp_request));
if let Some(rs) = relay_state {
headers.push_str(&ecp_relay_state_header_xml(rs));
}
Ok(crate::bindings::soap::soap_envelope_wrap(
authn_request_xml,
Some(&headers),
))
}
pub fn build_ecp_phase2_envelope(
saml_response_xml: &str,
paos_response: &PaosResponse,
relay_state: Option<&EcpRelayState>,
) -> Result<String, BindingError> {
let mut headers = String::new();
headers.push_str(&paos_response_header_xml(paos_response));
if let Some(rs) = relay_state {
headers.push_str(&ecp_relay_state_header_xml(rs));
}
Ok(crate::bindings::soap::soap_envelope_wrap(
saml_response_xml,
Some(&headers),
))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_paos_request_header() {
let req = PaosRequest {
response_consumer_url: "https://sp.example.com/acs".to_string(),
service: Some("My SP".to_string()),
message_id: Some("_msg001".to_string()),
};
let xml = paos_request_header_xml(&req);
assert!(xml.contains("paos:Request"));
assert!(xml.contains("mustUnderstand=\"1\""));
assert!(xml.contains("responseConsumerURL=\"https://sp.example.com/acs\""));
assert!(xml.contains("service=\"My SP\""));
assert!(xml.contains("messageID=\"_msg001\""));
}
#[test]
fn test_ecp_response_header() {
let resp = EcpResponse {
assertion_consumer_service_url: "https://sp.example.com/acs".to_string(),
};
let xml = ecp_response_header_xml(&resp);
assert!(xml.contains("ecp:Response"));
assert!(xml.contains("AssertionConsumerServiceURL="));
}
#[test]
fn test_ecp_relay_state_header() {
let rs = EcpRelayState {
relay_state: "token123".to_string(),
};
let xml = ecp_relay_state_header_xml(&rs);
assert!(xml.contains("ecp:RelayState"));
assert!(xml.contains("token123"));
}
#[test]
fn test_ecp_phase1_envelope() {
let authn_req = "<samlp:AuthnRequest/>";
let ecp_req = EcpRequest {
issuer: Some("https://sp.example.com".to_string()),
provider_name: Some("Test SP".to_string()),
is_passive: false,
idp_list: vec!["https://idp.example.com".to_string()],
};
let paos_req = PaosRequest {
response_consumer_url: "https://sp.example.com/acs".to_string(),
service: Some(ECP_NS.to_string()),
message_id: Some("_msg001".to_string()),
};
let env = build_ecp_phase1_envelope(authn_req, &ecp_req, &paos_req, None).unwrap();
assert!(env.contains("soap:Envelope"));
assert!(env.contains("soap:Header"));
assert!(env.contains("ecp:Request"));
assert!(env.contains("paos:Request"));
assert!(env.contains("saml:Issuer"));
assert!(env.contains("samlp:IDPEntry"));
assert!(env.contains("responseConsumerURL=\"https://sp.example.com/acs\""));
assert!(env.contains("AuthnRequest"));
}
#[test]
fn test_ecp_request_header_empty_children_self_closes() {
let req = EcpRequest {
issuer: None,
provider_name: None,
is_passive: true,
idp_list: vec![],
};
let xml = ecp_request_header_xml(&req);
assert!(xml.ends_with("/>"));
assert!(xml.contains("IsPassive=\"true\""));
}
#[test]
fn test_xml_escape_in_headers() {
let resp = EcpResponse {
assertion_consumer_service_url: "https://sp.example.com/acs?a=1&b=2".to_string(),
};
let xml = ecp_response_header_xml(&resp);
assert!(xml.contains("a=1&b=2"));
}
#[test]
fn test_ecp_phase2_envelope() {
let saml_resp = "<samlp:Response/>";
let paos_resp = PaosResponse {
ref_to_message_id: Some("_msg001".to_string()),
};
let rs = EcpRelayState {
relay_state: "abc".to_string(),
};
let env = build_ecp_phase2_envelope(saml_resp, &paos_resp, Some(&rs)).unwrap();
assert!(env.contains("soap:Envelope"));
assert!(env.contains("paos:Response"));
assert!(env.contains("ecp:RelayState"));
assert!(env.contains("abc"));
assert!(env.contains("Response"));
}
}