Skip to main content

http_msgsign_draft/sign/
input.rs

1use crate::errors::{HttpPayloadSeekError, InvalidValue, SignatureInputError, VerificationError};
2use crate::sign::{SignatureBase, VerifierKey};
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    pub fn created(&self) -> Option<u64> {
31        self.created
32    }
33
34    pub fn expires(&self) -> Option<u64> {
35        self.expires
36    }
37
38    pub fn verify_request<B>(
39        &self,
40        request: &Request<B>,
41        key: &impl VerifierKey
42    ) -> Result<(), VerificationError>
43    where
44        B: http_body::Body + Send,
45        B::Data: Send
46    {
47        let base = self.seek_request(request)?;
48        base.verify(key, &self.signature)?;
49        Ok(())
50    }
51    
52    pub fn verify_response<B>(
53        &self,
54        response: &Response<B>,
55        key: &impl VerifierKey
56    ) -> Result<(), VerificationError>
57    where
58        B: http_body::Body + Send,
59        B::Data: Send
60    {
61        let base = self.seek_response(response)?;
62        base.verify(key, &self.signature)?;
63        Ok(())
64    }
65}
66
67impl SignatureInput {
68    pub(crate) fn seek_request<B>(
69        &self,
70        request: &Request<B>,
71    ) -> Result<SignatureBase, HttpPayloadSeekError> {
72        let seeked = self
73            .headers
74            .iter()
75            .map(|target| target.seek_request(request))
76            .collect::<Result<IndexSet<_>, _>>()?;
77
78        Ok(SignatureBase::from_components(seeked))
79    }
80
81    pub(crate) fn seek_response<B>(
82        &self,
83        response: &Response<B>,
84    ) -> Result<SignatureBase, HttpPayloadSeekError> {
85        let seeked = self
86            .headers
87            .iter()
88            .map(|target| target.seek_response(response))
89            .collect::<Result<IndexSet<_>, _>>()?;
90
91        Ok(SignatureBase::from_components(seeked))
92    }
93
94    pub fn from_header(header: &HeaderMap) -> Result<SignatureInput, SignatureInputError> {
95        // Look for the Signature header defined in the RFC,
96        // or if none, look for the Authorization header and get its value.
97        // In the case of an Authorization header,
98        // remove the `Signature ` prefix and set the value to the same state as the value of the Signature header.
99        let params = match header.get("signature") {
100            Some(params) => params.to_str(),
101            None => match header.get(http::header::AUTHORIZATION) {
102                Some(params) => params
103                    .to_str()
104                    .map(|params| params.strip_prefix("Signature ").unwrap()),
105                None => return Err(SignatureInputError::NotExist),
106            },
107        }
108        .map_err(|_| InvalidValue::String)?;
109
110        Self::parse(params)
111    }
112
113    fn parse(value: &str) -> Result<SignatureInput, SignatureInputError> {
114        let params = value
115            .split(',')
116            .flat_map(|st| st.split_once('='))
117            .map(|(key, value)| (key, value.trim_matches('"')))
118            .collect::<HashMap<&str, &str>>();
119
120        let Some(key_id) = params.get("keyId") else {
121            return Err(SignatureInputError::RequireParameter("keyId"));
122        };
123
124        // This is not required within the document (also RECOMMENDED) but would be almost mandatory.
125        let Some(algorithm) = params.get("algorithm") else {
126            return Err(SignatureInputError::RequireParameter("algorithm"));
127        };
128
129        let created = params
130            .get("created")
131            .map(|s| u64::from_str(s))
132            .transpose()
133            .map_err(|_| InvalidValue::Integer)?;
134
135        let expires = params
136            .get("expires")
137            .map(|s| u64::from_str(s))
138            .transpose()
139            .map_err(|_| InvalidValue::Integer)?;
140
141        let headers = match params.get("headers") {
142            None => {
143                if created.is_none() {
144                    return Err(SignatureInputError::RequireParameter("created"));
145                }
146
147                // If not specified, implementations MUST operate as if the field were specified
148                // with a single value, `(created)`, in the list of HTTP headers.
149                indexset! {
150                    TargetField::Created(created)
151                }
152            }
153            Some(headers) => headers
154                .split(' ')
155                .map(|field| {
156                    Ok(match field {
157                        REQUEST_TARGET => TargetField::RequestTarget,
158                        CREATED => {
159                            if created.is_none() {
160                                return Err(SignatureInputError::RequireParameter("created"));
161                            }
162                            TargetField::Created(created)
163                        }
164                        EXPIRES => TargetField::Expires(TimeOrDuration::Time(
165                            expires.ok_or(SignatureInputError::RequireParameter("expires"))?,
166                        )),
167                        header => TargetField::HeaderField(HeaderName::from_str(header).unwrap()),
168                    })
169                })
170                .collect::<Result<IndexSet<TargetField>, _>>()?,
171        };
172
173        // A zero-length `headers` parameter value MUST NOT be used.
174        if headers.is_empty() {
175            return Err(InvalidValue::NonEmptyArray)?;
176        }
177
178        let Some(signature) = params
179            .get("signature")
180            .map(|encoded| base64::engine::general_purpose::STANDARD.decode(encoded))
181            .transpose()
182            .unwrap()
183        else {
184            return Err(SignatureInputError::RequireParameter("signature"));
185        };
186
187        Ok(Self {
188            key_id: key_id.to_string(),
189            algorithm: algorithm.to_string(),
190            created,
191            expires,
192            headers,
193            signature,
194        })
195    }
196}
197
198impl<B> TryFrom<&Request<B>> for SignatureInput
199where
200    B: http_body::Body + Send,
201    B::Data: Send
202{
203    type Error = SignatureInputError;
204    
205    fn try_from(value: &Request<B>) -> Result<Self, Self::Error> {
206        Self::from_header(value.headers())
207    }
208}
209
210impl<B> TryFrom<&Response<B>> for SignatureInput
211where
212    B: http_body::Body + Send,
213    B::Data: Send
214{
215    type Error = SignatureInputError;
216    
217    fn try_from(value: &Response<B>) -> Result<Self, Self::Error> {
218        Self::from_header(value.headers())
219    }
220}