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
9pub(super) fn stretch_key(key: &Pin<Box<Array<u8, U32>>>) -> Aes256CbcHmacKey {
13 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
21pub(crate) fn pad_bytes(bytes: &mut Vec<u8>, min_length: usize) -> Result<(), CryptoError> {
26 let pad_bytes = min_length.saturating_sub(bytes.len()).max(1);
28 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
38pub(crate) fn unpad_bytes(padded_bytes: &[u8]) -> Result<&[u8], CryptoError> {
42 let pad_len = *padded_bytes.last().ok_or(CryptoError::InvalidPadding)? as usize;
43 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 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 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}