1#![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 let mut n_found = 0;
92 #[allow(clippy::needless_range_loop)]
93 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 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 assert!(SystemParams::generate() == SystemParams::get_hardcoded());
153
154 let key_pair = KeyPair::derive_from(sho.as_mut());
155
156 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 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}