Skip to main content

saml_rs/browser/
input.rs

1//! Inbound browser message conversion helpers.
2//!
3//! References: SAML Bindings 2.0 <https://docs.oasis-open.org/security/saml/v2.0/saml-bindings-2.0-os.pdf> and HTTP POST-SimpleSign <https://docs.oasis-open.org/security/saml/Post2.0/sstc-saml-binding-simplesign.html>.
4
5use core::marker::PhantomData;
6
7use super::forms::FormField;
8use crate::binding::{base64_decode_with_limit, build_simplesign_octet};
9use crate::constants::url_params;
10use crate::error::SamlError;
11use crate::model::{AuthnRequest, LogoutRequest, LogoutResponse, SsoResponse};
12use crate::raw::HttpRequest;
13use crate::xml::XmlLimits;
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq)]
16enum MessageField {
17    Request,
18    Response,
19}
20
21impl MessageField {
22    fn param_name(self) -> &'static str {
23        match self {
24            Self::Request => url_params::SAML_REQUEST,
25            Self::Response => url_params::SAML_RESPONSE,
26        }
27    }
28
29    fn opposite_param_name(self) -> &'static str {
30        match self {
31            Self::Request => url_params::SAML_RESPONSE,
32            Self::Response => url_params::SAML_REQUEST,
33        }
34    }
35}
36
37/// Typed browser input for inbound SAML messages.
38///
39/// # Examples
40///
41/// ```
42/// use saml_rs::{AuthnRequest, BrowserInput, FormField, SsoResponse};
43///
44/// let redirect = BrowserInput::<AuthnRequest>::redirect("SAMLRequest=...");
45/// let post = BrowserInput::<SsoResponse>::post(vec![
46///     FormField::new("SAMLResponse", "..."),
47/// ]);
48///
49/// # let _ = (redirect, post);
50/// ```
51///
52/// SSO responses are received through POST-family bindings, not Redirect:
53///
54/// ```compile_fail
55/// use saml_rs::{BrowserInput, SsoResponse};
56///
57/// let _ = BrowserInput::<SsoResponse>::redirect("SAMLResponse=...");
58/// ```
59#[derive(Debug, Clone, PartialEq, Eq)]
60pub enum BrowserInput<Message> {
61    /// HTTP-Redirect input as a raw query string.
62    Redirect {
63        /// Raw URL query, with or without a leading `?`.
64        raw_query: String,
65        /// Message marker.
66        _message: PhantomData<Message>,
67    },
68    /// HTTP-POST input as parsed form fields.
69    Post {
70        /// Parsed form fields.
71        fields: Vec<FormField>,
72        /// Message marker.
73        _message: PhantomData<Message>,
74    },
75    /// HTTP-POST-SimpleSign input.
76    SimpleSignPost {
77        /// Raw form body.
78        raw_body: String,
79        /// Parsed form fields.
80        fields: Vec<FormField>,
81        /// Message marker.
82        _message: PhantomData<Message>,
83    },
84}
85
86impl BrowserInput<AuthnRequest> {
87    /// Create Redirect input from a raw query string.
88    pub fn redirect(raw_query: impl Into<String>) -> Self {
89        redirect_input(raw_query)
90    }
91
92    /// Create POST input from parsed fields.
93    pub fn post(fields: Vec<FormField>) -> Self {
94        post_input(fields)
95    }
96
97    /// Create SimpleSign input from parsed fields.
98    pub fn simple_sign(fields: Vec<FormField>) -> Self {
99        simple_sign_input(fields)
100    }
101
102    /// Parse a raw `application/x-www-form-urlencoded` SimpleSign body.
103    pub fn simple_sign_body(raw_body: impl Into<String>) -> Self {
104        simple_sign_body_input(raw_body)
105    }
106}
107
108impl BrowserInput<SsoResponse> {
109    /// Create POST input from parsed fields.
110    pub fn post(fields: Vec<FormField>) -> Self {
111        post_input(fields)
112    }
113
114    /// Create SimpleSign input from parsed fields.
115    pub fn simple_sign(fields: Vec<FormField>) -> Self {
116        simple_sign_input(fields)
117    }
118
119    /// Parse a raw `application/x-www-form-urlencoded` SimpleSign body.
120    pub fn simple_sign_body(raw_body: impl Into<String>) -> Self {
121        simple_sign_body_input(raw_body)
122    }
123}
124
125impl BrowserInput<LogoutRequest> {
126    /// Create Redirect input from a raw query string.
127    pub fn redirect(raw_query: impl Into<String>) -> Self {
128        redirect_input(raw_query)
129    }
130
131    /// Create POST input from parsed fields.
132    pub fn post(fields: Vec<FormField>) -> Self {
133        post_input(fields)
134    }
135
136    /// Create SimpleSign input from parsed fields.
137    pub fn simple_sign(fields: Vec<FormField>) -> Self {
138        simple_sign_input(fields)
139    }
140
141    /// Parse a raw `application/x-www-form-urlencoded` SimpleSign body.
142    pub fn simple_sign_body(raw_body: impl Into<String>) -> Self {
143        simple_sign_body_input(raw_body)
144    }
145}
146
147impl BrowserInput<LogoutResponse> {
148    /// Create Redirect input from a raw query string.
149    pub fn redirect(raw_query: impl Into<String>) -> Self {
150        redirect_input(raw_query)
151    }
152
153    /// Create POST input from parsed fields.
154    pub fn post(fields: Vec<FormField>) -> Self {
155        post_input(fields)
156    }
157
158    /// Create SimpleSign input from parsed fields.
159    pub fn simple_sign(fields: Vec<FormField>) -> Self {
160        simple_sign_input(fields)
161    }
162
163    /// Parse a raw `application/x-www-form-urlencoded` SimpleSign body.
164    pub fn simple_sign_body(raw_body: impl Into<String>) -> Self {
165        simple_sign_body_input(raw_body)
166    }
167}
168
169impl TryFrom<BrowserInput<AuthnRequest>> for HttpRequest {
170    type Error = SamlError;
171
172    fn try_from(value: BrowserInput<AuthnRequest>) -> Result<Self, Self::Error> {
173        http_request_from_input(value, true, MessageField::Request)
174    }
175}
176
177impl TryFrom<BrowserInput<SsoResponse>> for HttpRequest {
178    type Error = SamlError;
179
180    fn try_from(value: BrowserInput<SsoResponse>) -> Result<Self, Self::Error> {
181        http_request_from_input(value, false, MessageField::Response)
182    }
183}
184
185impl TryFrom<BrowserInput<LogoutRequest>> for HttpRequest {
186    type Error = SamlError;
187
188    fn try_from(value: BrowserInput<LogoutRequest>) -> Result<Self, Self::Error> {
189        http_request_from_input(value, true, MessageField::Request)
190    }
191}
192
193impl TryFrom<BrowserInput<LogoutResponse>> for HttpRequest {
194    type Error = SamlError;
195
196    fn try_from(value: BrowserInput<LogoutResponse>) -> Result<Self, Self::Error> {
197        http_request_from_input(value, true, MessageField::Response)
198    }
199}
200
201fn redirect_input<Message>(raw_query: impl Into<String>) -> BrowserInput<Message> {
202    BrowserInput::Redirect {
203        raw_query: raw_query.into(),
204        _message: PhantomData,
205    }
206}
207
208fn post_input<Message>(fields: Vec<FormField>) -> BrowserInput<Message> {
209    BrowserInput::Post {
210        fields,
211        _message: PhantomData,
212    }
213}
214
215fn simple_sign_input<Message>(fields: Vec<FormField>) -> BrowserInput<Message> {
216    BrowserInput::SimpleSignPost {
217        raw_body: String::new(),
218        fields,
219        _message: PhantomData,
220    }
221}
222
223fn simple_sign_body_input<Message>(raw_body: impl Into<String>) -> BrowserInput<Message> {
224    let raw_body = raw_body.into();
225    let fields = parse_form_fields(&raw_body);
226    BrowserInput::SimpleSignPost {
227        raw_body,
228        fields,
229        _message: PhantomData,
230    }
231}
232
233fn http_request_from_input<Message>(
234    value: BrowserInput<Message>,
235    allow_redirect: bool,
236    expected_message: MessageField,
237) -> Result<HttpRequest, SamlError> {
238    match value {
239        BrowserInput::Redirect { raw_query, .. } => {
240            if !allow_redirect {
241                return Err(SamlError::UndefinedBinding);
242            }
243            let raw_query = raw_query.trim_start_matches('?').to_string();
244            let octet_string = redirect_octet_from_raw_query(&raw_query, expected_message)?;
245            let query = parse_form_pairs(&raw_query);
246            Ok(HttpRequest {
247                query,
248                octet_string,
249                ..Default::default()
250            })
251        }
252        BrowserInput::Post { fields, .. } => {
253            validate_form_message_kind(&fields, expected_message)?;
254            Ok(HttpRequest::post(fields_to_pairs(fields)))
255        }
256        BrowserInput::SimpleSignPost {
257            raw_body,
258            mut fields,
259            ..
260        } => {
261            if fields.is_empty() {
262                fields = parse_form_fields(&raw_body);
263            }
264            let octet_string = simplesign_octet_from_fields(&fields, expected_message)?;
265            let mut request = HttpRequest::post(fields_to_pairs(fields));
266            request.octet_string = Some(octet_string);
267            Ok(request)
268        }
269    }
270}
271
272struct RawQueryParam<'a> {
273    name: &'a str,
274    encoded_value: &'a str,
275}
276
277fn raw_query_params(raw_query: &str) -> impl Iterator<Item = RawQueryParam<'_>> {
278    raw_query
279        .split('&')
280        .filter(|segment| !segment.is_empty())
281        .map(|segment| match segment.split_once('=') {
282            Some((name, encoded_value)) => RawQueryParam {
283                name,
284                encoded_value,
285            },
286            None => RawQueryParam {
287                name: segment,
288                encoded_value: "",
289            },
290        })
291}
292
293fn unique_encoded_query_value<'a>(
294    params: &'a [RawQueryParam<'a>],
295    name: &str,
296) -> Result<Option<&'a str>, SamlError> {
297    let mut values = params
298        .iter()
299        .filter(|param| param.name == name)
300        .map(|param| param.encoded_value);
301    let first = values.next();
302    if values.next().is_some() {
303        return Err(SamlError::Invalid(format!(
304            "ambiguous Redirect field {name}"
305        )));
306    }
307    Ok(first)
308}
309
310fn validate_query_message_kind<'a>(
311    params: &'a [RawQueryParam<'a>],
312    expected_message: MessageField,
313) -> Result<&'a str, SamlError> {
314    let expected_name = expected_message.param_name();
315    let opposite_name = expected_message.opposite_param_name();
316    let expected = unique_encoded_query_value(params, expected_name)?;
317    let opposite = unique_encoded_query_value(params, opposite_name)?;
318    match (expected, opposite) {
319        (Some(value), None) => Ok(value),
320        (None, None) => Err(SamlError::Invalid(format!(
321            "missing Redirect field {expected_name}"
322        ))),
323        (None, Some(_)) | (Some(_), Some(_)) => Err(SamlError::Invalid(format!(
324            "expected Redirect field {expected_name}, found {opposite_name}"
325        ))),
326    }
327}
328
329fn redirect_octet_from_raw_query(
330    raw_query: &str,
331    expected_message: MessageField,
332) -> Result<Option<String>, SamlError> {
333    let params: Vec<_> = raw_query_params(raw_query).collect();
334    let message_name = expected_message.param_name();
335    let message_value = validate_query_message_kind(&params, expected_message)?;
336    let relay_state = unique_encoded_query_value(&params, url_params::RELAY_STATE)?;
337    let sig_alg = unique_encoded_query_value(&params, url_params::SIG_ALG)?;
338    let signature = unique_encoded_query_value(&params, url_params::SIGNATURE)?;
339
340    let sig_alg = match (sig_alg, signature) {
341        (Some(sig_alg), Some(_)) => sig_alg,
342        (None, None) => return Ok(None),
343        _ => {
344            return Err(SamlError::Invalid(
345                "incomplete Redirect signature parameters".into(),
346            ))
347        }
348    };
349
350    Ok(Some(match relay_state {
351        Some(relay_state) => {
352            format!("{message_name}={message_value}&RelayState={relay_state}&SigAlg={sig_alg}")
353        }
354        None => format!("{message_name}={message_value}&SigAlg={sig_alg}"),
355    }))
356}
357
358fn parse_form_pairs(raw: &str) -> Vec<(String, String)> {
359    url::form_urlencoded::parse(raw.as_bytes())
360        .map(|(name, value)| (name.into_owned(), value.into_owned()))
361        .collect()
362}
363
364fn parse_form_fields(raw: &str) -> Vec<FormField> {
365    parse_form_pairs(raw)
366        .into_iter()
367        .map(|(name, value)| FormField::new(name, value))
368        .collect()
369}
370
371fn fields_to_pairs(fields: Vec<FormField>) -> Vec<(String, String)> {
372    fields.into_iter().map(FormField::into_pair).collect()
373}
374
375fn field_value<'a>(fields: &'a [FormField], name: &str) -> Result<Option<&'a str>, SamlError> {
376    let mut values = fields
377        .iter()
378        .filter(|field| field.name() == name)
379        .map(FormField::value);
380    let first = values.next();
381    if values.next().is_some() {
382        return Err(SamlError::Invalid("ambiguous form field".into()));
383    }
384    Ok(first)
385}
386
387fn validate_form_message_kind(
388    fields: &[FormField],
389    expected_message: MessageField,
390) -> Result<&str, SamlError> {
391    let expected_name = expected_message.param_name();
392    let opposite_name = expected_message.opposite_param_name();
393    let expected = field_value(fields, expected_name)?;
394    let opposite = field_value(fields, opposite_name)?;
395    match (expected, opposite) {
396        (Some(value), None) => Ok(value),
397        (None, None) => Err(SamlError::Invalid(format!(
398            "missing form field {expected_name}"
399        ))),
400        (None, Some(_)) | (Some(_), Some(_)) => Err(SamlError::Invalid(format!(
401            "expected form field {expected_name}, found {opposite_name}"
402        ))),
403    }
404}
405
406fn simplesign_octet_from_fields(
407    fields: &[FormField],
408    expected_message: MessageField,
409) -> Result<String, SamlError> {
410    let message_name = expected_message.param_name();
411    let encoded = validate_form_message_kind(fields, expected_message)?;
412    let raw_xml = String::from_utf8(base64_decode_with_limit(
413        encoded,
414        XmlLimits::default().max_bytes,
415    )?)
416    .map_err(|err| SamlError::Xml(err.to_string()))?;
417    let sig_alg = field_value(fields, url_params::SIG_ALG)?.ok_or(SamlError::MissingSigAlg)?;
418    let relay_state = field_value(fields, url_params::RELAY_STATE)?;
419    Ok(build_simplesign_octet(
420        message_name,
421        &raw_xml,
422        relay_state,
423        sig_alg,
424    ))
425}