1use 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#[derive(Debug, Clone, PartialEq, Eq)]
60pub enum BrowserInput<Message> {
61 Redirect {
63 raw_query: String,
65 _message: PhantomData<Message>,
67 },
68 Post {
70 fields: Vec<FormField>,
72 _message: PhantomData<Message>,
74 },
75 SimpleSignPost {
77 raw_body: String,
79 fields: Vec<FormField>,
81 _message: PhantomData<Message>,
83 },
84}
85
86impl BrowserInput<AuthnRequest> {
87 pub fn redirect(raw_query: impl Into<String>) -> Self {
89 redirect_input(raw_query)
90 }
91
92 pub fn post(fields: Vec<FormField>) -> Self {
94 post_input(fields)
95 }
96
97 pub fn simple_sign(fields: Vec<FormField>) -> Self {
99 simple_sign_input(fields)
100 }
101
102 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 pub fn post(fields: Vec<FormField>) -> Self {
111 post_input(fields)
112 }
113
114 pub fn simple_sign(fields: Vec<FormField>) -> Self {
116 simple_sign_input(fields)
117 }
118
119 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 pub fn redirect(raw_query: impl Into<String>) -> Self {
128 redirect_input(raw_query)
129 }
130
131 pub fn post(fields: Vec<FormField>) -> Self {
133 post_input(fields)
134 }
135
136 pub fn simple_sign(fields: Vec<FormField>) -> Self {
138 simple_sign_input(fields)
139 }
140
141 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 pub fn redirect(raw_query: impl Into<String>) -> Self {
150 redirect_input(raw_query)
151 }
152
153 pub fn post(fields: Vec<FormField>) -> Self {
155 post_input(fields)
156 }
157
158 pub fn simple_sign(fields: Vec<FormField>) -> Self {
160 simple_sign_input(fields)
161 }
162
163 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(¶ms, expected_message)?;
336 let relay_state = unique_encoded_query_value(¶ms, url_params::RELAY_STATE)?;
337 let sig_alg = unique_encoded_query_value(¶ms, url_params::SIG_ALG)?;
338 let signature = unique_encoded_query_value(¶ms, 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}