http_msgsign_draft/sign/
input.rs

1use crate::errors::{HttpPayloadSeekError, InvalidValue, SignatureInputError};
2use crate::sign::SignatureBase;
3use crate::sign::field::{CREATED, EXPIRES, REQUEST_TARGET, TargetField, TimeOrDuration};
4use base64::Engine;
5use http::{HeaderMap, HeaderName, Request, Response};
6use indexmap::{IndexSet, indexset};
7use std::collections::HashMap;
8use std::str::FromStr;
9
10#[allow(unused)]
11#[derive(Debug, Clone)]
12pub struct SignatureInput {
13    pub(crate) key_id: String,
14    pub(crate) algorithm: String,
15    pub(crate) created: Option<u64>,
16    pub(crate) expires: Option<u64>,
17    pub(crate) headers: IndexSet<TargetField>,
18    pub(crate) signature: Vec<u8>,
19}
20
21impl SignatureInput {
22    pub fn key_id(&self) -> &str {
23        &self.key_id
24    }
25
26    pub fn algorithm(&self) -> &str {
27        &self.algorithm
28    }
29}
30
31impl SignatureInput {
32    pub(crate) fn seek_request<B>(
33        &self,
34        request: &Request<B>,
35    ) -> Result<SignatureBase, HttpPayloadSeekError> {
36        let seeked = self
37            .headers
38            .iter()
39            .map(|target| target.seek_request(request))
40            .collect::<Result<IndexSet<_>, _>>()?;
41
42        Ok(SignatureBase::from_components(seeked))
43    }
44
45    pub(crate) fn seek_response<B>(
46        &self,
47        response: &Response<B>,
48    ) -> Result<SignatureBase, HttpPayloadSeekError> {
49        let seeked = self
50            .headers
51            .iter()
52            .map(|target| target.seek_response(response))
53            .collect::<Result<IndexSet<_>, _>>()?;
54
55        Ok(SignatureBase::from_components(seeked))
56    }
57
58    pub fn from_header(header: &HeaderMap) -> Result<SignatureInput, SignatureInputError> {
59        // Look for the Signature header defined in the RFC,
60        // or if none, look for the Authorization header and get its value.
61        // In the case of an Authorization header,
62        // remove the `Signature ` prefix and set the value to the same state as the value of the Signature header.
63        let params = match header.get("signature") {
64            Some(params) => params.to_str(),
65            None => match header.get(http::header::AUTHORIZATION) {
66                Some(params) => params
67                    .to_str()
68                    .map(|params| params.strip_prefix("Signature ").unwrap()),
69                None => return Err(SignatureInputError::NotExist),
70            },
71        }
72        .map_err(|_| InvalidValue::String)?;
73
74        Self::parse(params)
75    }
76
77    fn parse(value: &str) -> Result<SignatureInput, SignatureInputError> {
78        let params = value
79            .split(',')
80            .flat_map(|st| st.split_once('='))
81            .map(|(key, value)| (key, value.trim_matches('"')))
82            .collect::<HashMap<&str, &str>>();
83
84        let Some(key_id) = params.get("keyId") else {
85            return Err(SignatureInputError::RequireParameter("keyId"));
86        };
87
88        // This is not required within the document (also RECOMMENDED) but would be almost mandatory.
89        let Some(algorithm) = params.get("algorithm") else {
90            return Err(SignatureInputError::RequireParameter("algorithm"));
91        };
92
93        let created = params
94            .get("created")
95            .map(|s| u64::from_str(s))
96            .transpose()
97            .map_err(|_| InvalidValue::Integer)?;
98
99        let expires = params
100            .get("expires")
101            .map(|s| u64::from_str(s))
102            .transpose()
103            .map_err(|_| InvalidValue::Integer)?;
104
105        let headers = match params.get("headers") {
106            None => {
107                if created.is_none() {
108                    return Err(SignatureInputError::RequireParameter("created"));
109                }
110
111                // If not specified, implementations MUST operate as if the field were specified
112                // with a single value, `(created)`, in the list of HTTP headers.
113                indexset! {
114                    TargetField::Created(created)
115                }
116            }
117            Some(headers) => headers
118                .split(' ')
119                .map(|field| {
120                    Ok(match field {
121                        REQUEST_TARGET => TargetField::RequestTarget,
122                        CREATED => {
123                            if created.is_none() {
124                                return Err(SignatureInputError::RequireParameter("created"));
125                            }
126                            TargetField::Created(created)
127                        }
128                        EXPIRES => TargetField::Expires(TimeOrDuration::Time(
129                            expires.ok_or(SignatureInputError::RequireParameter("expires"))?,
130                        )),
131                        header => TargetField::HeaderField(HeaderName::from_str(header).unwrap()),
132                    })
133                })
134                .collect::<Result<IndexSet<TargetField>, _>>()?,
135        };
136
137        // A zero-length `headers` parameter value MUST NOT be used.
138        if headers.is_empty() {
139            return Err(InvalidValue::NonEmptyArray)?;
140        }
141
142        let Some(signature) = params
143            .get("signature")
144            .map(|encoded| base64::engine::general_purpose::STANDARD.decode(encoded))
145            .transpose()
146            .unwrap()
147        else {
148            return Err(SignatureInputError::RequireParameter("signature"));
149        };
150
151        Ok(Self {
152            key_id: key_id.to_string(),
153            algorithm: algorithm.to_string(),
154            created,
155            expires,
156            headers,
157            signature,
158        })
159    }
160}