Skip to main content

wasi_crypto/symmetric/
state.rs

1use std::convert::TryFrom;
2use std::sync::{Arc, Mutex, MutexGuard};
3
4use super::*;
5use crate::CryptoCtx;
6use crate::Limits;
7
8#[derive(Clone)]
9pub struct SymmetricState {
10    inner: Arc<Mutex<Box<dyn SymmetricStateLike>>>,
11}
12
13impl SymmetricState {
14    fn new(symmetric_state_like: Box<dyn SymmetricStateLike>) -> Self {
15        SymmetricState {
16            inner: Arc::new(Mutex::new(symmetric_state_like)),
17        }
18    }
19
20    fn inner(&self) -> MutexGuard<'_, Box<dyn SymmetricStateLike>> {
21        self.inner.lock().unwrap()
22    }
23
24    fn locked<T, U>(&self, mut f: T) -> U
25    where
26        T: FnMut(MutexGuard<'_, Box<dyn SymmetricStateLike>>) -> U,
27    {
28        f(self.inner())
29    }
30
31    fn open(
32        alg_str: &str,
33        key: Option<SymmetricKey>,
34        options: Option<SymmetricOptions>,
35        limits: &Limits,
36    ) -> Result<SymmetricState, CryptoError> {
37        let alg = SymmetricAlgorithm::try_from(alg_str)?;
38        if let Some(ref key) = key {
39            ensure!(key.alg() == alg, CryptoError::InvalidKey);
40        }
41        let size_limit = limits.0.get(alg_str).copied();
42        let symmetric_state = match alg {
43            SymmetricAlgorithm::HmacSha256 | SymmetricAlgorithm::HmacSha512 => SymmetricState::new(
44                Box::new(HmacSha2SymmetricState::new(alg, key, options, size_limit)?),
45            ),
46            SymmetricAlgorithm::Sha256
47            | SymmetricAlgorithm::Sha384
48            | SymmetricAlgorithm::Sha512
49            | SymmetricAlgorithm::Sha512_256 => SymmetricState::new(Box::new(
50                Sha2SymmetricState::new(alg, None, options, size_limit)?,
51            )),
52            SymmetricAlgorithm::HkdfSha256Expand
53            | SymmetricAlgorithm::HkdfSha512Expand
54            | SymmetricAlgorithm::HkdfSha256Extract
55            | SymmetricAlgorithm::HkdfSha512Extract => SymmetricState::new(Box::new(
56                HkdfSymmetricState::new(alg, key, options, size_limit)?,
57            )),
58            SymmetricAlgorithm::Aes128Gcm | SymmetricAlgorithm::Aes256Gcm => SymmetricState::new(
59                Box::new(AesGcmSymmetricState::new(alg, key, options, size_limit)?),
60            ),
61            SymmetricAlgorithm::ChaCha20Poly1305 => SymmetricState::new(Box::new(
62                ChaChaPolySymmetricState::new(alg, key, options, size_limit)?,
63            )),
64            SymmetricAlgorithm::XChaCha20Poly1305 => SymmetricState::new(Box::new(
65                ChaChaPolySymmetricState::new(alg, key, options, size_limit)?,
66            )),
67            SymmetricAlgorithm::Xoodyak128 | SymmetricAlgorithm::Xoodyak160 => SymmetricState::new(
68                Box::new(XoodyakSymmetricState::new(alg, key, options, size_limit)?),
69            ),
70            _ => bail!(CryptoError::UnsupportedAlgorithm),
71        };
72        Ok(symmetric_state)
73    }
74}
75
76pub trait SymmetricStateLike: Sync + Send {
77    fn alg(&self) -> SymmetricAlgorithm;
78    fn options_get(&self, name: &str) -> Result<Vec<u8>, CryptoError>;
79    fn options_get_u64(&self, name: &str) -> Result<u64, CryptoError>;
80
81    fn size_limit(&self) -> Option<usize>;
82
83    fn absorb_unchecked(&mut self, _data: &[u8]) -> Result<(), CryptoError> {
84        bail!(CryptoError::InvalidOperation)
85    }
86
87    fn absorb(&mut self, data: &[u8]) -> Result<(), CryptoError> {
88        ensure!(
89            self.size_limit().is_none_or(|l| data.len() <= l),
90            CryptoError::Overflow
91        );
92        self.absorb_unchecked(data)
93    }
94
95    fn squeeze_unchecked(&mut self, _out: &mut [u8]) -> Result<(), CryptoError> {
96        bail!(CryptoError::InvalidOperation)
97    }
98
99    fn squeeze(&mut self, out: &mut [u8]) -> Result<(), CryptoError> {
100        ensure!(
101            self.size_limit().is_none_or(|l| out.len() <= l),
102            CryptoError::Overflow
103        );
104        self.squeeze_unchecked(out)
105    }
106
107    fn squeeze_key(&mut self, _alg_str: &str) -> Result<SymmetricKey, CryptoError> {
108        bail!(CryptoError::InvalidOperation)
109    }
110
111    fn squeeze_tag(&mut self) -> Result<SymmetricTag, CryptoError> {
112        bail!(CryptoError::InvalidOperation)
113    }
114
115    fn max_tag_len(&mut self) -> Result<usize, CryptoError> {
116        bail!(CryptoError::InvalidOperation)
117    }
118
119    fn encrypt_unchecked(&mut self, _out: &mut [u8], _data: &[u8]) -> Result<usize, CryptoError> {
120        bail!(CryptoError::InvalidOperation)
121    }
122
123    fn encrypt(&mut self, out: &mut [u8], data: &[u8]) -> Result<usize, CryptoError> {
124        ensure!(
125            out.len()
126                == data
127                    .len()
128                    .checked_add(self.max_tag_len()?)
129                    .ok_or(CryptoError::Overflow)?,
130            CryptoError::InvalidLength
131        );
132        ensure!(
133            self.size_limit().is_none_or(|l| data.len() <= l),
134            CryptoError::Overflow
135        );
136        self.encrypt_unchecked(out, data)
137    }
138
139    fn encrypt_detached_unchecked(
140        &mut self,
141        _out: &mut [u8],
142        _data: &[u8],
143    ) -> Result<SymmetricTag, CryptoError> {
144        bail!(CryptoError::InvalidOperation)
145    }
146
147    fn encrypt_detached(
148        &mut self,
149        out: &mut [u8],
150        data: &[u8],
151    ) -> Result<SymmetricTag, CryptoError> {
152        ensure!(out.len() == data.len(), CryptoError::InvalidLength);
153        ensure!(
154            self.size_limit().is_none_or(|l| data.len() <= l),
155            CryptoError::Overflow
156        );
157        self.encrypt_detached_unchecked(out, data)
158    }
159
160    fn decrypt_unchecked(&mut self, _out: &mut [u8], _data: &[u8]) -> Result<usize, CryptoError> {
161        bail!(CryptoError::InvalidOperation)
162    }
163
164    fn decrypt(&mut self, out: &mut [u8], data: &[u8]) -> Result<usize, CryptoError> {
165        ensure!(
166            out.len()
167                == data
168                    .len()
169                    .checked_sub(self.max_tag_len()?)
170                    .ok_or(CryptoError::Overflow)?,
171            CryptoError::Overflow
172        );
173        ensure!(
174            self.size_limit().is_none_or(|l| data.len() <= l),
175            CryptoError::Overflow
176        );
177        match self.decrypt_unchecked(out, data) {
178            Ok(out_len) => Ok(out_len),
179            Err(e) => {
180                out.iter_mut().for_each(|x| *x = 0);
181                Err(e)
182            }
183        }
184    }
185
186    fn decrypt_detached_unchecked(
187        &mut self,
188        _out: &mut [u8],
189        _data: &[u8],
190        _raw_tag: &[u8],
191    ) -> Result<usize, CryptoError> {
192        bail!(CryptoError::InvalidOperation)
193    }
194
195    fn decrypt_detached(
196        &mut self,
197        out: &mut [u8],
198        data: &[u8],
199        raw_tag: &[u8],
200    ) -> Result<usize, CryptoError> {
201        ensure!(out.len() == data.len(), CryptoError::InvalidLength);
202        ensure!(
203            self.size_limit().is_none_or(|l| data.len() <= l),
204            CryptoError::Overflow
205        );
206        match self.decrypt_detached_unchecked(out, data, raw_tag) {
207            Ok(out_len) => Ok(out_len),
208            Err(e) => {
209                out.iter_mut().for_each(|x| *x = 0);
210                Err(e)
211            }
212        }
213    }
214
215    fn ratchet(&mut self) -> Result<(), CryptoError> {
216        bail!(CryptoError::InvalidOperation)
217    }
218}
219
220impl CryptoCtx {
221    pub fn symmetric_state_open(
222        &self,
223        alg_str: &str,
224        key_handle: Option<Handle>,
225        options_handle: Option<Handle>,
226    ) -> Result<Handle, CryptoError> {
227        let key = match key_handle {
228            None => None,
229            Some(symmetric_key_handle) => {
230                Some(self.handles.symmetric_key.get(symmetric_key_handle)?)
231            }
232        };
233        let options = match options_handle {
234            None => None,
235            Some(options_handle) => {
236                Some(self.handles.options.get(options_handle)?.into_symmetric()?)
237            }
238        };
239        let symmetric_state = SymmetricState::open(alg_str, key, options, &self.limits)?;
240        let handle = self.handles.symmetric_state.register(symmetric_state)?;
241        Ok(handle)
242    }
243
244    pub fn symmetric_state_options_get(
245        &self,
246        symmetric_state_handle: Handle,
247        name: &str,
248        value: &mut [u8],
249    ) -> Result<usize, CryptoError> {
250        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
251        let v = symmetric_state.inner().options_get(name)?;
252        let v_len = v.len();
253        ensure!(v_len <= value.len(), CryptoError::Overflow);
254        value[..v_len].copy_from_slice(&v);
255        Ok(v_len)
256    }
257
258    pub fn symmetric_state_options_get_u64(
259        &self,
260        symmetric_state_handle: Handle,
261        name: &str,
262    ) -> Result<u64, CryptoError> {
263        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
264        let v = symmetric_state.inner().options_get_u64(name)?;
265        Ok(v)
266    }
267
268    pub fn symmetric_state_close(&self, symmetric_state_handle: Handle) -> Result<(), CryptoError> {
269        self.handles.symmetric_state.close(symmetric_state_handle)
270    }
271
272    pub fn symmetric_state_clone(
273        &self,
274        symmetric_state_handle: Handle,
275    ) -> Result<Handle, CryptoError> {
276        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
277        let symmetric_state = symmetric_state.clone();
278        let handle = self.handles.symmetric_state.register(symmetric_state)?;
279        Ok(handle)
280    }
281
282    pub fn symmetric_state_absorb(
283        &self,
284        symmetric_state_handle: Handle,
285        data: &[u8],
286    ) -> Result<(), CryptoError> {
287        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
288        symmetric_state.locked(|mut state| state.absorb_unchecked(data))
289    }
290
291    pub fn symmetric_state_squeeze(
292        &self,
293        symmetric_state_handle: Handle,
294        out: &mut [u8],
295    ) -> Result<(), CryptoError> {
296        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
297        symmetric_state.locked(|mut state| state.squeeze_unchecked(out))
298    }
299
300    pub fn symmetric_state_squeeze_tag(
301        &self,
302        symmetric_state_handle: Handle,
303    ) -> Result<Handle, CryptoError> {
304        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
305        let tag = symmetric_state.locked(|mut state| state.squeeze_tag())?;
306        let handle = self.handles.symmetric_tag.register(tag)?;
307        Ok(handle)
308    }
309
310    pub fn symmetric_state_squeeze_key(
311        &self,
312        symmetric_state_handle: Handle,
313        alg_str: &str,
314    ) -> Result<Handle, CryptoError> {
315        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
316        let symmetric_key = symmetric_state.locked(|mut state| state.squeeze_key(alg_str))?;
317        let handle = self.handles.symmetric_key.register(symmetric_key)?;
318        Ok(handle)
319    }
320
321    pub fn symmetric_state_max_tag_len(
322        &self,
323        symmetric_state_handle: Handle,
324    ) -> Result<usize, CryptoError> {
325        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
326        let max_tag_len = symmetric_state.inner().max_tag_len()?;
327        Ok(max_tag_len)
328    }
329
330    pub fn symmetric_state_encrypt(
331        &self,
332        symmetric_state_handle: Handle,
333        out: &mut [u8],
334        data: &[u8],
335    ) -> Result<usize, CryptoError> {
336        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
337        symmetric_state.locked(|mut state| state.encrypt(out, data))
338    }
339
340    pub fn symmetric_state_encrypt_detached(
341        &self,
342        symmetric_state_handle: Handle,
343        out: &mut [u8],
344        data: &[u8],
345    ) -> Result<Handle, CryptoError> {
346        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
347        let symmetric_tag = symmetric_state.inner().encrypt_detached(out, data)?;
348        let handle = self.handles.symmetric_tag.register(symmetric_tag)?;
349        Ok(handle)
350    }
351
352    pub fn symmetric_state_decrypt(
353        &self,
354        symmetric_state_handle: Handle,
355        out: &mut [u8],
356        data: &[u8],
357    ) -> Result<usize, CryptoError> {
358        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
359        symmetric_state.locked(|mut state| state.decrypt(out, data))
360    }
361
362    pub fn symmetric_state_decrypt_detached(
363        &self,
364        symmetric_state_handle: Handle,
365        out: &mut [u8],
366        data: &[u8],
367        raw_tag: &[u8],
368    ) -> Result<usize, CryptoError> {
369        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
370        symmetric_state.locked(|mut state| state.decrypt_detached(out, data, raw_tag))
371    }
372
373    pub fn symmetric_state_ratchet(
374        &self,
375        symmetric_state_handle: Handle,
376    ) -> Result<(), CryptoError> {
377        let symmetric_state = self.handles.symmetric_state.get(symmetric_state_handle)?;
378        symmetric_state.locked(|mut state| state.ratchet())
379    }
380}