Skip to main content

libsignal_service/
profile_cipher.rs

1use std::convert::TryInto;
2
3use aes_gcm::{aead::Aead, AeadInOut, Aes256Gcm, KeyInit};
4use rand::{CryptoRng, RngCore};
5use zkgroup::profiles::ProfileKey;
6
7use crate::{
8    profile_name::ProfileName, websocket::profile::SignalServiceProfile,
9    Profile,
10};
11
12/// Encrypt and decrypt a [`ProfileName`] and other profile information.
13///
14/// # Example
15///
16/// ```rust
17/// # use libsignal_service::{profile_name::ProfileName, profile_cipher::ProfileCipher};
18/// # use zkgroup::profiles::ProfileKey;
19/// # use rand::Rng;
20/// # let mut rng = rand::rng();
21/// # let some_randomness = rng.random();
22/// let profile_key = ProfileKey::generate(some_randomness);
23/// let name = ProfileName::<&str> {
24///     given_name: "Bill",
25///     family_name: None,
26/// };
27/// let cipher = ProfileCipher::new(profile_key);
28/// let encrypted = cipher.encrypt_name(&name, &mut rng).unwrap();
29/// let decrypted = cipher.decrypt_name(&encrypted).unwrap().unwrap();
30/// assert_eq!(decrypted.as_ref(), name);
31/// ```
32pub struct ProfileCipher {
33    profile_key: ProfileKey,
34}
35
36const NAME_PADDED_LENGTH_1: usize = 53;
37const NAME_PADDED_LENGTH_2: usize = 257;
38const NAME_PADDING_BRACKETS: &[usize] =
39    &[NAME_PADDED_LENGTH_1, NAME_PADDED_LENGTH_2];
40
41const ABOUT_PADDED_LENGTH_1: usize = 128;
42const ABOUT_PADDED_LENGTH_2: usize = 254;
43const ABOUT_PADDED_LENGTH_3: usize = 512;
44const ABOUT_PADDING_BRACKETS: &[usize] = &[
45    ABOUT_PADDED_LENGTH_1,
46    ABOUT_PADDED_LENGTH_2,
47    ABOUT_PADDED_LENGTH_3,
48];
49
50const EMOJI_PADDED_LENGTH: usize = 32;
51
52#[derive(thiserror::Error, Debug)]
53pub enum ProfileCipherError {
54    #[error("Encryption error")]
55    EncryptionError,
56    #[error("UTF-8 decode error {0}")]
57    Utf8Error(#[from] std::str::Utf8Error),
58    #[error("Input name too long")]
59    InputTooLong,
60}
61
62fn pad_plaintext(
63    bytes: &mut Vec<u8>,
64    brackets: &[usize],
65) -> Result<usize, ProfileCipherError> {
66    let len = brackets
67        .iter()
68        .find(|x| **x >= bytes.len())
69        .ok_or(ProfileCipherError::InputTooLong)?;
70    let len: usize = *len;
71
72    bytes.resize(len, 0);
73    assert!(brackets.contains(&bytes.len()));
74
75    Ok(len)
76}
77
78impl ProfileCipher {
79    pub fn new(profile_key: ProfileKey) -> Self {
80        Self { profile_key }
81    }
82
83    pub fn into_inner(self) -> ProfileKey {
84        self.profile_key
85    }
86
87    fn pad_and_encrypt<R: RngCore + CryptoRng>(
88        &self,
89        mut bytes: Vec<u8>,
90        padding_brackets: &[usize],
91        csprng: &mut R,
92    ) -> Result<Vec<u8>, ProfileCipherError> {
93        let _len = pad_plaintext(&mut bytes, padding_brackets)?;
94
95        let cipher = Aes256Gcm::new(&self.profile_key.get_bytes().into());
96        let mut nonce = [0u8; 12];
97        csprng.fill_bytes(&mut nonce);
98
99        cipher
100            .encrypt_in_place(&nonce.into(), b"", &mut bytes)
101            .map_err(|_| ProfileCipherError::EncryptionError)?;
102
103        let mut concat = Vec::with_capacity(nonce.len() + bytes.len());
104        concat.extend_from_slice(&nonce);
105        concat.extend_from_slice(&bytes);
106        Ok(concat)
107    }
108
109    fn decrypt_and_unpad(
110        &self,
111        bytes: impl AsRef<[u8]>,
112    ) -> Result<Vec<u8>, ProfileCipherError> {
113        let bytes = bytes.as_ref();
114        let nonce: [u8; 12] = bytes[0..12]
115            .try_into()
116            .expect("fixed length nonce material");
117        let cipher = Aes256Gcm::new(&self.profile_key.get_bytes().into());
118
119        let mut plaintext = cipher
120            .decrypt(&nonce.into(), &bytes[12..])
121            .map_err(|_| ProfileCipherError::EncryptionError)?;
122
123        // Unpad
124        let len = plaintext
125            .iter()
126            // Search the first non-0 char...
127            .rposition(|x| *x != 0)
128            // ...and strip until right after.
129            .map(|x| x + 1)
130            // If it's all zeroes, the string is 0-length.
131            .unwrap_or(0);
132        plaintext.truncate(len);
133        Ok(plaintext)
134    }
135
136    pub fn decrypt(
137        &self,
138        encrypted_profile: SignalServiceProfile,
139    ) -> Result<Profile, ProfileCipherError> {
140        let name = encrypted_profile
141            .name
142            .as_ref()
143            .map(|data| self.decrypt_name(data))
144            .transpose()?
145            .flatten();
146        let about = encrypted_profile
147            .about
148            .as_ref()
149            .map(|data| self.decrypt_about(data))
150            .transpose()?;
151        let about_emoji = encrypted_profile
152            .about_emoji
153            .as_ref()
154            .map(|data| self.decrypt_emoji(data))
155            .transpose()?;
156
157        Ok(Profile {
158            name,
159            about,
160            about_emoji,
161            avatar: encrypted_profile.avatar,
162            unrestricted_unidentified_access: encrypted_profile
163                .unrestricted_unidentified_access,
164        })
165    }
166
167    pub fn decrypt_avatar(
168        &self,
169        bytes: &[u8],
170    ) -> Result<Vec<u8>, ProfileCipherError> {
171        self.decrypt_and_unpad(bytes)
172    }
173
174    pub fn encrypt_name<'inp, R: RngCore + CryptoRng>(
175        &self,
176        name: impl std::borrow::Borrow<ProfileName<&'inp str>>,
177        csprng: &mut R,
178    ) -> Result<Vec<u8>, ProfileCipherError> {
179        let name = name.borrow();
180        let bytes = name.serialize();
181        self.pad_and_encrypt(bytes, NAME_PADDING_BRACKETS, csprng)
182    }
183
184    pub fn decrypt_name(
185        &self,
186        bytes: impl AsRef<[u8]>,
187    ) -> Result<Option<ProfileName<String>>, ProfileCipherError> {
188        let bytes = bytes.as_ref();
189
190        let plaintext = self.decrypt_and_unpad(bytes)?;
191
192        Ok(ProfileName::<String>::deserialize(&plaintext)?)
193    }
194
195    pub fn encrypt_about<R: RngCore + CryptoRng>(
196        &self,
197        about: String,
198        csprng: &mut R,
199    ) -> Result<Vec<u8>, ProfileCipherError> {
200        let bytes = about.into_bytes();
201        self.pad_and_encrypt(bytes, ABOUT_PADDING_BRACKETS, csprng)
202    }
203
204    pub fn decrypt_about(
205        &self,
206        bytes: impl AsRef<[u8]>,
207    ) -> Result<String, ProfileCipherError> {
208        let bytes = bytes.as_ref();
209
210        let plaintext = self.decrypt_and_unpad(bytes)?;
211
212        // XXX This re-allocates.
213        Ok(std::str::from_utf8(&plaintext)?.into())
214    }
215
216    pub fn encrypt_emoji<R: RngCore + CryptoRng>(
217        &self,
218        emoji: String,
219        csprng: &mut R,
220    ) -> Result<Vec<u8>, ProfileCipherError> {
221        let bytes = emoji.into_bytes();
222        self.pad_and_encrypt(bytes, &[EMOJI_PADDED_LENGTH], csprng)
223    }
224
225    pub fn decrypt_emoji(
226        &self,
227        bytes: impl AsRef<[u8]>,
228    ) -> Result<String, ProfileCipherError> {
229        let bytes = bytes.as_ref();
230
231        let plaintext = self.decrypt_and_unpad(bytes)?;
232
233        // XXX This re-allocates.
234        Ok(std::str::from_utf8(&plaintext)?.into())
235    }
236}
237
238#[cfg(test)]
239mod tests {
240    use super::*;
241    use crate::profile_name::ProfileName;
242    use rand::Rng;
243    use zkgroup::profiles::ProfileKey;
244
245    #[test]
246    fn roundtrip_name() {
247        let names = [
248            "Me and my guitar", // shorter that 53
249            "abcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyz", // one shorter than 53
250            "abcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzx", // exactly 53
251            "abcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzxf", // one more than 53
252            "abcdefghijklmnopqrstuvwxyzabcdefghijklmnopqrstuvwxyzxfoobar", // a bit more than 53
253        ];
254
255        // Test the test cases
256        assert_eq!(names[1].len(), NAME_PADDED_LENGTH_1 - 1);
257        assert_eq!(names[2].len(), NAME_PADDED_LENGTH_1);
258        assert_eq!(names[3].len(), NAME_PADDED_LENGTH_1 + 1);
259
260        let mut rng = rand::rng();
261        let some_randomness = rng.random();
262        let profile_key = ProfileKey::generate(some_randomness);
263        let cipher = ProfileCipher::new(profile_key);
264        for name in &names {
265            let profile_name = ProfileName::<&str> {
266                given_name: name,
267                family_name: None,
268            };
269            assert_eq!(profile_name.serialize().len(), name.len());
270            let encrypted =
271                cipher.encrypt_name(&profile_name, &mut rng).unwrap();
272            let decrypted = cipher.decrypt_name(encrypted).unwrap().unwrap();
273
274            assert_eq!(decrypted.as_ref(), profile_name);
275            assert_eq!(decrypted.serialize(), profile_name.serialize());
276            assert_eq!(&decrypted.given_name, name);
277        }
278    }
279
280    #[test]
281    fn roundtrip_about() {
282        let abouts = [
283            "Me and my guitar", // shorter that 53
284        ];
285
286        let mut rng = rand::rng();
287        let some_randomness = rng.random();
288        let profile_key = ProfileKey::generate(some_randomness);
289        let cipher = ProfileCipher::new(profile_key);
290
291        for &about in &abouts {
292            let encrypted =
293                cipher.encrypt_about(about.into(), &mut rng).unwrap();
294            let decrypted = cipher.decrypt_about(encrypted).unwrap();
295
296            assert_eq!(decrypted, about);
297        }
298    }
299
300    #[test]
301    fn roundtrip_emoji() {
302        let emojii = ["❤️", "💩", "🤣", "😲", "🐠"];
303
304        let mut rng = rand::rng();
305        let some_randomness = rng.random();
306        let profile_key = ProfileKey::generate(some_randomness);
307        let cipher = ProfileCipher::new(profile_key);
308
309        for &emoji in &emojii {
310            let encrypted =
311                cipher.encrypt_emoji(emoji.into(), &mut rng).unwrap();
312            let decrypted = cipher.decrypt_emoji(encrypted).unwrap();
313
314            assert_eq!(decrypted, emoji);
315        }
316    }
317}