Skip to main content

zkgroup/crypto/
profile_key_encryption.rs

1//
2// Copyright 2020 Signal Messenger, LLC.
3// SPDX-License-Identifier: AGPL-3.0-only
4//
5
6#![allow(non_snake_case)]
7
8use std::sync::LazyLock;
9
10use curve25519_dalek::ristretto::RistrettoPoint;
11use partial_default::PartialDefault;
12use serde::{Deserialize, Serialize};
13use subtle::{Choice, ConstantTimeEq, CtOption};
14use zkcredential::attributes::Attribute;
15
16use crate::common::errors::*;
17use crate::common::sho::*;
18use crate::common::simple_types::*;
19use crate::crypto::profile_key_struct;
20
21static SYSTEM_PARAMS: LazyLock<SystemParams> = LazyLock::new(|| {
22    crate::deserialize(&SystemParams::SYSTEM_HARDCODED).expect("valid hardcoded params")
23});
24
25#[derive(Copy, Clone, PartialEq, Eq, Serialize, Deserialize, PartialDefault)]
26pub struct SystemParams {
27    pub(crate) G_b1: RistrettoPoint,
28    pub(crate) G_b2: RistrettoPoint,
29}
30
31pub type KeyPair = zkcredential::attributes::KeyPair<ProfileKeyEncryptionDomain>;
32pub type PublicKey = zkcredential::attributes::PublicKey<ProfileKeyEncryptionDomain>;
33pub type Ciphertext = zkcredential::attributes::Ciphertext<ProfileKeyEncryptionDomain>;
34
35impl SystemParams {
36    pub fn generate() -> Self {
37        let mut sho = Sho::new(
38            b"Signal_ZKGroup_20200424_Constant_ProfileKeyEncryption_SystemParams_Generate",
39            b"",
40        );
41        let G_b1 = sho.get_point();
42        let G_b2 = sho.get_point();
43        SystemParams { G_b1, G_b2 }
44    }
45
46    pub fn get_hardcoded() -> SystemParams {
47        *SYSTEM_PARAMS
48    }
49
50    const SYSTEM_HARDCODED: [u8; 64] = [
51        0xf6, 0xba, 0xa3, 0x17, 0xce, 0x18, 0x39, 0xc9, 0x3d, 0x61, 0x7e, 0xc, 0xd8, 0x37, 0xd1,
52        0x9d, 0xa9, 0xc8, 0xa4, 0xc5, 0x20, 0xbf, 0x7c, 0x51, 0xb1, 0xe6, 0xc2, 0xcb, 0x2a, 0x4,
53        0x9c, 0x61, 0x2e, 0x1, 0x75, 0x89, 0x4c, 0x87, 0x30, 0xb2, 0x3, 0xab, 0x3b, 0xd9, 0x8e,
54        0xcb, 0x2d, 0x81, 0xab, 0xac, 0xb6, 0x5f, 0x8a, 0x61, 0x24, 0xf4, 0x97, 0x71, 0xd1, 0x4a,
55        0x98, 0x52, 0x12, 0xc,
56    ];
57}
58
59pub struct ProfileKeyEncryptionDomain;
60impl zkcredential::attributes::Domain for ProfileKeyEncryptionDomain {
61    type Attribute = profile_key_struct::ProfileKeyStruct;
62
63    const ID: &'static str = "Signal_ZKGroup_20231011_ProfileKeyEncryption";
64
65    fn G_a() -> [RistrettoPoint; 2] {
66        let system = SystemParams::get_hardcoded();
67        [system.G_b1, system.G_b2]
68    }
69}
70
71impl ProfileKeyEncryptionDomain {
72    pub(crate) fn decrypt(
73        key_pair: &KeyPair,
74        ciphertext: &Ciphertext,
75        uid_bytes: UidBytes,
76    ) -> Result<profile_key_struct::ProfileKeyStruct, ZkGroupVerificationFailure> {
77        let M4 = key_pair
78            .decrypt_to_second_point(ciphertext)
79            .map_err(|_| ZkGroupVerificationFailure)?;
80        let candidates = M4.map_to_curve_inverse();
81
82        let target_M3 = key_pair.a1.invert() * ciphertext.as_points()[0];
83        let seed_sho = profile_key_struct::ProfileKeyStruct::seed_M3();
84
85        let mut retval = CtOption::new(
86            profile_key_struct::ProfileKeyStruct::partial_default(),
87            Choice::from(0u8),
88        );
89        // The number of valid solutions found. Note we would see if n_found > 1 because the closures in
90        // CtOption::and_then and CtOption::or_else always run.
91        let mut n_found = 0;
92        #[allow(clippy::needless_range_loop)]
93        // Only iterate the first 8 solutions, i.e., the positive ones.
94        for i in 0..8 {
95            retval = retval.or_else(|| {
96                candidates[i].and_then(|profile_key_bytes| {
97                    let mut candidate_retval = CtOption::new(
98                        profile_key_struct::ProfileKeyStruct::partial_default(),
99                        Choice::from(0u8),
100                    );
101                    for j in 0..8 {
102                        candidate_retval = candidate_retval.or_else(|| {
103                            let mut pk = profile_key_bytes;
104                            if ((j >> 2) & 1) == 1 {
105                                pk[0] |= 0x01;
106                            }
107                            if ((j >> 1) & 1) == 1 {
108                                pk[31] |= 0x80;
109                            }
110                            if (j & 1) == 1 {
111                                pk[31] |= 0x40;
112                            }
113                            let M3 = profile_key_struct::ProfileKeyStruct::calc_M3(
114                                seed_sho.clone(),
115                                pk,
116                                uid_bytes,
117                            );
118                            let found = M3.ct_eq(&target_M3);
119                            n_found += found.unwrap_u8();
120                            CtOption::new(
121                                profile_key_struct::ProfileKeyStruct { bytes: pk, M3, M4 },
122                                found,
123                            )
124                        });
125                    }
126                    candidate_retval
127                })
128            });
129        }
130        if n_found == 1 {
131            // We can unwrap because n_found > 0 implies that candidate_retval is Some, which means
132            // retval is Some
133            Ok(retval.unwrap())
134        } else {
135            Err(ZkGroupVerificationFailure)
136        }
137    }
138}
139
140#[cfg(test)]
141mod tests {
142    use super::*;
143    use crate::common::constants::*;
144
145    #[test]
146    fn test_profile_key_encryption() {
147        let master_key = TEST_ARRAY_32_1;
148        let mut sho = Sho::new(b"Test_Profile_Key_Encryption", &master_key);
149
150        //let system = SystemParams::generate();
151        //println!("PARAMS = {:#x?}", bincode::serialize(&system));
152        assert!(SystemParams::generate() == SystemParams::get_hardcoded());
153
154        let key_pair = KeyPair::derive_from(sho.as_mut());
155
156        // Test serialize of key_pair
157        let key_pair_bytes = bincode::serialize(&key_pair).unwrap();
158        match bincode::deserialize::<KeyPair>(&key_pair_bytes[0..key_pair_bytes.len() - 1]) {
159            Err(_) => (),
160            _ => unreachable!(),
161        };
162        let key_pair2: KeyPair = bincode::deserialize(&key_pair_bytes).unwrap();
163        assert!(key_pair == key_pair2);
164
165        let profile_key_bytes = TEST_ARRAY_32_1;
166        let uid_bytes = TEST_ARRAY_16_1;
167        let profile_key = profile_key_struct::ProfileKeyStruct::new(profile_key_bytes, uid_bytes);
168        let ciphertext = key_pair.encrypt(&profile_key);
169
170        // Test serialize / deserialize of Ciphertext
171        let ciphertext_bytes = bincode::serialize(&ciphertext).unwrap();
172        assert!(ciphertext_bytes.len() == 64);
173        let ciphertext2: Ciphertext = bincode::deserialize(&ciphertext_bytes).unwrap();
174        assert!(ciphertext == ciphertext2);
175        println!("ciphertext_bytes = {ciphertext_bytes:#x?}");
176        assert!(
177            ciphertext_bytes
178                == vec![
179                    0x56, 0x18, 0xcb, 0x4c, 0x7d, 0x72, 0x1e, 0x1, 0x2b, 0x22, 0xf0, 0x77, 0xef,
180                    0x12, 0x64, 0xf6, 0xb1, 0x43, 0xbb, 0x59, 0x7a, 0x1d, 0x66, 0x5a, 0x70, 0xaa,
181                    0x84, 0x24, 0x5f, 0x24, 0x6d, 0x20, 0xba, 0xdb, 0x97, 0x47, 0x4a, 0x56, 0xf4,
182                    0xb5, 0x36, 0x1a, 0xec, 0xa9, 0xd1, 0x18, 0xb7, 0x0, 0x4e, 0x14, 0x9, 0x71,
183                    0x99, 0xa, 0xab, 0x2a, 0xf2, 0x43, 0x2d, 0x3f, 0x8f, 0x7d, 0x21, 0x3a,
184                ]
185        );
186
187        let plaintext =
188            ProfileKeyEncryptionDomain::decrypt(&key_pair, &ciphertext2, uid_bytes).unwrap();
189        assert!(plaintext == profile_key);
190
191        let mut sho = Sho::new(b"Test_Repeated_ProfileKeyEnc/Dec", b"seed");
192        for _ in 0..100 {
193            let uid_bytes: UidBytes = sho.squeeze_as_array();
194            let profile_key_bytes: ProfileKeyBytes = sho.squeeze_as_array();
195
196            let profile_key =
197                profile_key_struct::ProfileKeyStruct::new(profile_key_bytes, uid_bytes);
198            let ciphertext = key_pair.encrypt(&profile_key);
199            assert!(
200                ProfileKeyEncryptionDomain::decrypt(&key_pair, &ciphertext, uid_bytes).unwrap()
201                    == profile_key
202            );
203        }
204
205        let uid_bytes = TEST_ARRAY_16;
206        let profile_key = profile_key_struct::ProfileKeyStruct::new(TEST_ARRAY_32, TEST_ARRAY_16);
207        let ciphertext = key_pair.encrypt(&profile_key);
208        assert!(
209            ProfileKeyEncryptionDomain::decrypt(&key_pair, &ciphertext, uid_bytes).unwrap()
210                == profile_key
211        );
212
213        let uid_bytes = TEST_ARRAY_16;
214        let profile_key = profile_key_struct::ProfileKeyStruct::new(TEST_ARRAY_32_2, TEST_ARRAY_16);
215        let ciphertext = key_pair.encrypt(&profile_key);
216        assert!(
217            ProfileKeyEncryptionDomain::decrypt(&key_pair, &ciphertext, uid_bytes).unwrap()
218                == profile_key
219        );
220
221        let uid_bytes = TEST_ARRAY_16;
222        let profile_key = profile_key_struct::ProfileKeyStruct::new(TEST_ARRAY_32_3, TEST_ARRAY_16);
223        let ciphertext = key_pair.encrypt(&profile_key);
224        assert!(
225            ProfileKeyEncryptionDomain::decrypt(&key_pair, &ciphertext, uid_bytes).unwrap()
226                == profile_key
227        );
228
229        let uid_bytes = TEST_ARRAY_16;
230        let profile_key = profile_key_struct::ProfileKeyStruct::new(TEST_ARRAY_32_4, TEST_ARRAY_16);
231        let ciphertext = key_pair.encrypt(&profile_key);
232        assert!(
233            ProfileKeyEncryptionDomain::decrypt(&key_pair, &ciphertext, uid_bytes).unwrap()
234                == profile_key
235        );
236    }
237}