Skip to main content

bitwarden_crypto/keys/
utils.rs

1use std::{cmp::max, pin::Pin};
2
3use hybrid_array::Array;
4use typenum::U32;
5
6use super::Aes256CbcHmacKey;
7use crate::{CryptoError, Result, util::hkdf_expand};
8
9/// Stretch the given key using HKDF.
10/// This can be either a kdf-derived key (PIN/Master password) or
11/// a random key from key connector
12pub(super) fn stretch_key(key: &Pin<Box<Array<u8, U32>>>) -> Aes256CbcHmacKey {
13    // this is safe because the key length is always 32 bytes
14    let enc_key: Pin<Box<Array<u8, U32>>> =
15        hkdf_expand(key, Some("enc")).expect("HKDF expand to succeed");
16    let mac_key: Pin<Box<Array<u8, U32>>> =
17        hkdf_expand(key, Some("mac")).expect("HKDF expand to succeed");
18    Aes256CbcHmacKey::new(&enc_key.0, &mac_key.0)
19}
20
21/// Pads bytes to a minimum length using PKCS7-like padding.
22/// The last N bytes of the padded bytes all have the value N. Minimum of 1 padding byte. maximum of
23/// 255 bytes. For example, padded to size 4, the value 0,0 becomes 0,0,2,2.
24/// If the more than 255 bytes of padding is required, this will return an error.
25pub(crate) fn pad_bytes(bytes: &mut Vec<u8>, min_length: usize) -> Result<(), CryptoError> {
26    // at least 1 byte of padding is required
27    let pad_bytes = min_length.saturating_sub(bytes.len()).max(1);
28    // Since each byte represents the padding size, the maximum padding is 255.
29    // If more padding is required, this will return an error.
30    if pad_bytes > 255 {
31        return Err(CryptoError::InvalidPadding);
32    }
33    let padded_length = max(min_length, bytes.len() + 1);
34    bytes.resize(padded_length, pad_bytes as u8);
35    Ok(())
36}
37
38/// Unpads bytes that is padded using the PKCS7-like padding defined by [pad_bytes].
39/// The last N bytes of the padded bytes all have the value N. Minimum of 1 padding byte.
40/// For example, padded to size 4, the value 0,0 becomes 0,0,2,2.
41pub(crate) fn unpad_bytes(padded_bytes: &[u8]) -> Result<&[u8], CryptoError> {
42    let pad_len = *padded_bytes.last().ok_or(CryptoError::InvalidPadding)? as usize;
43    // The padding is at minimum 1 as noted in `pad_bytes`
44    if pad_len == 0 || pad_len > padded_bytes.len() {
45        return Err(CryptoError::InvalidPadding);
46    }
47    Ok(padded_bytes[..(padded_bytes.len() - pad_len)].as_ref())
48}
49
50#[cfg(test)]
51mod tests {
52    use super::*;
53
54    #[test]
55    fn test_stretch_kdf_key() {
56        let key = Box::pin(
57            [
58                31, 79, 104, 226, 150, 71, 177, 90, 194, 80, 172, 209, 17, 129, 132, 81, 138, 167,
59                69, 167, 254, 149, 2, 27, 39, 197, 64, 42, 22, 195, 86, 75,
60            ]
61            .into(),
62        );
63        let stretched = stretch_key(&key);
64
65        assert_eq!(
66            [
67                111, 31, 178, 45, 238, 152, 37, 114, 143, 215, 124, 83, 135, 173, 195, 23, 142,
68                134, 120, 249, 61, 132, 163, 182, 113, 197, 189, 204, 188, 21, 237, 96
69            ],
70            stretched.enc_key().as_slice()
71        );
72        assert_eq!(
73            [
74                221, 127, 206, 234, 101, 27, 202, 38, 86, 52, 34, 28, 78, 28, 185, 16, 48, 61, 127,
75                166, 209, 247, 194, 87, 232, 26, 48, 85, 193, 249, 179, 155
76            ],
77            stretched.mac_key().as_slice()
78        );
79    }
80
81    #[test]
82    fn test_pad_bytes_256_error() {
83        let mut bytes = vec![1u8; 0];
84        let result = pad_bytes(&mut bytes, 256);
85        assert!(matches!(result, Err(CryptoError::InvalidPadding)));
86    }
87
88    #[test]
89    fn test_pad_bytes_roundtrip() {
90        let original_bytes = vec![1u8; 10];
91        let mut cloned_bytes = original_bytes.clone();
92        let mut encoded_bytes = vec![1u8; 12];
93        encoded_bytes[10] = 2;
94        encoded_bytes[11] = 2;
95        pad_bytes(&mut cloned_bytes, 12).expect("Padding failed");
96        assert_eq!(encoded_bytes, cloned_bytes);
97        let unpadded_bytes = unpad_bytes(&cloned_bytes).unwrap();
98        assert_eq!(original_bytes, unpadded_bytes);
99    }
100
101    #[test]
102    fn test_pad_bytes_roundtrip_empty() {
103        let original_bytes = Vec::new();
104        let mut cloned_bytes = original_bytes.clone();
105        pad_bytes(&mut cloned_bytes, 32).expect("Padding failed");
106        let unpadded = unpad_bytes(&cloned_bytes).unwrap();
107        assert_eq!(Vec::<u8>::new(), unpadded);
108    }
109
110    #[test]
111    fn test_unpad_bytes_invalid_empty() {
112        let data: Vec<u8> = vec![];
113        let result = unpad_bytes(&data);
114        assert!(matches!(result, Err(CryptoError::InvalidPadding)));
115    }
116
117    #[test]
118    fn test_unpad_bytes_invalid_too_large() {
119        // Last byte is 5, but only 4 bytes in total
120        let data = vec![1, 2, 3, 5];
121        let result = unpad_bytes(&data);
122        assert!(matches!(result, Err(CryptoError::InvalidPadding)));
123    }
124
125    #[test]
126    fn test_unpad_bytes_invalid_0_padding() {
127        // Padding value of 0 is invalid
128        let data = vec![1, 2, 3, 0];
129        let result = unpad_bytes(&data);
130        assert!(matches!(result, Err(CryptoError::InvalidPadding)));
131    }
132
133    #[test]
134    fn test_pad_and_unpad_bytes_range_0_to_1024() {
135        let cases: Vec<_> = (0..=1024)
136            .flat_map(|data_size| (2..=1024).map(move |padding_size| (data_size, padding_size)))
137            .collect();
138
139        let data_larger_than_padding_cases: Vec<_> = cases
140            .clone()
141            .into_iter()
142            .filter(|(data_size, padding_size)| data_size > padding_size)
143            .collect();
144        for (data_size, padding_size) in data_larger_than_padding_cases {
145            let mut data: Vec<u8> = vec![0x12; data_size];
146            let original = data.clone();
147            pad_bytes(&mut data, padding_size).expect("Padding failed");
148            let unpadded = unpad_bytes(&data).expect("Unpadding failed");
149            assert_eq!(
150                unpadded, original,
151                "Failed at size {} and padding {}",
152                data_size, padding_size
153            );
154        }
155
156        let padding_larger_than_data_cases: Vec<_> = cases
157            .clone()
158            .into_iter()
159            .filter(|(data_size, padding_size)| {
160                data_size <= padding_size && (padding_size - data_size) <= 255
161            })
162            .collect();
163        for (data_size, padding_size) in padding_larger_than_data_cases {
164            println!(
165                "Testing data_size: {}, padding_size: {}",
166                data_size, padding_size
167            );
168            let data_original: Vec<u8> = vec![0x12; data_size];
169            let mut data = data_original.clone();
170
171            pad_bytes(&mut data, padding_size).expect("Padding failed");
172            let unpadded = unpad_bytes(&data).expect("Unpadding failed");
173            assert_eq!(
174                unpadded, data_original,
175                "Failed at size {} and padding {}",
176                data_size, padding_size
177            );
178        }
179
180        let padding_massively_larger_than_data_cases: Vec<_> = cases
181            .into_iter()
182            .filter(|(data_size, padding_size)| {
183                data_size <= padding_size && (padding_size - data_size) > 255
184            })
185            .collect();
186        for (data_size, padding_size) in padding_massively_larger_than_data_cases {
187            let mut data: Vec<u8> = vec![0x12; data_size];
188            let error = pad_bytes(&mut data, padding_size);
189            assert!(
190                matches!(error, Err(CryptoError::InvalidPadding)),
191                "Expected InvalidPadding error at size {} and padding {}, but got {:?}",
192                data_size,
193                padding_size,
194                error
195            );
196        }
197    }
198}