1use crate::ecies::{PrivateKey, PublicKey, RecoveryPackage, AES_KEY_LENGTH};
5use crate::nizk::{DLNizk, DdhTupleNizk};
6use crate::random_oracle::RandomOracle;
7use fastcrypto::aes::{Aes256Ctr, AesKey, Cipher, InitializationVector};
8use fastcrypto::error::{FastCryptoError, FastCryptoResult};
9use fastcrypto::groups::{FiatShamirChallenge, GroupElement, Scalar};
10use fastcrypto::hmac::{hkdf_sha3_256, HkdfIkm};
11use fastcrypto::traits::{AllowedRng, ToFromBytes};
12use serde::de::DeserializeOwned;
13use serde::{Deserialize, Serialize};
14use typenum::consts::{U16, U32};
15
16#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
29pub struct Encryption<G: GroupElement> {
30 ephemeral_key: G,
31 data: Vec<u8>,
32 hkdf_info: usize,
33}
34
35#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
38pub struct MultiRecipientEncryption<G: GroupElement>(G, Vec<Vec<u8>>, DLNizk<G>);
39
40impl<G> PrivateKey<G>
41where
42 G: GroupElement + Serialize,
43 <G as GroupElement>::ScalarType: FiatShamirChallenge,
44{
45 pub fn new<R: AllowedRng>(rng: &mut R) -> Self {
46 Self(G::ScalarType::rand(rng))
47 }
48
49 pub fn from(sc: G::ScalarType) -> Self {
50 Self(sc)
51 }
52
53 pub fn decrypt(&self, enc: &Encryption<G>) -> Vec<u8> {
54 enc.decrypt(&self.0)
55 }
56
57 pub fn create_recovery_package<R: AllowedRng>(
58 &self,
59 enc: &Encryption<G>,
60 random_oracle: &RandomOracle,
61 rng: &mut R,
62 ) -> RecoveryPackage<G> {
63 let ephemeral_key = enc.ephemeral_key * self.0;
64 let pk = G::generator() * self.0;
65 let proof = DdhTupleNizk::<G>::create(
66 &self.0,
67 &enc.ephemeral_key,
68 &pk,
69 &ephemeral_key,
70 random_oracle,
71 rng,
72 );
73 RecoveryPackage {
74 ephemeral_key,
75 proof,
76 }
77 }
78}
79
80impl<G> PublicKey<G>
81where
82 G: GroupElement + Serialize + DeserializeOwned,
83 <G as GroupElement>::ScalarType: FiatShamirChallenge,
84{
85 pub fn from_private_key(sk: &PrivateKey<G>) -> Self {
86 Self(G::generator() * sk.0)
87 }
88
89 #[cfg(test)]
90 pub fn encrypt<R: AllowedRng>(&self, msg: &[u8], rng: &mut R) -> Encryption<G> {
91 Encryption::<G>::encrypt(&self.0, msg, rng)
92 }
93
94 pub fn deterministic_encrypt(msg: &[u8], r_g: &G, r_x_g: &G, info: usize) -> Encryption<G> {
95 Encryption::<G>::deterministic_encrypt(msg, r_g, r_x_g, info)
96 }
97
98 pub fn decrypt_with_recovery_package(
99 &self,
100 pkg: &RecoveryPackage<G>,
101 random_oracle: &RandomOracle,
102 enc: &Encryption<G>,
103 ) -> FastCryptoResult<Vec<u8>> {
104 pkg.proof.verify(
105 &enc.ephemeral_key,
106 &self.0,
107 &pkg.ephemeral_key,
108 random_oracle,
109 )?;
110 Ok(enc.decrypt_from_partial_decryption(&pkg.ephemeral_key))
111 }
112
113 pub fn as_element(&self) -> &G {
114 &self.0
115 }
116}
117
118impl<G: GroupElement> From<G> for PublicKey<G> {
119 fn from(p: G) -> Self {
120 Self(p)
121 }
122}
123
124impl<G: GroupElement + Serialize> Encryption<G> {
125 fn sym_encrypt(k: &G, info: usize) -> Aes256Ctr {
126 Aes256Ctr::new(
127 AesKey::<U32>::from_bytes(&Self::hkdf(k, info))
128 .expect("New shouldn't fail as use fixed size key is used"),
129 )
130 }
131 fn deterministic_encrypt(msg: &[u8], r_g: &G, r_x_g: &G, hkdf_info: usize) -> Self {
132 let cipher = Self::sym_encrypt(r_x_g, hkdf_info);
133 let data = cipher.encrypt(&Self::fixed_zero_nonce(), msg);
134 Self {
135 ephemeral_key: *r_g,
136 data,
137 hkdf_info,
138 }
139 }
140
141 #[cfg(test)]
142 fn encrypt<R: AllowedRng>(x_g: &G, msg: &[u8], rng: &mut R) -> Self {
143 let r = G::ScalarType::rand(rng);
144 let r_g = G::generator() * r;
145 let r_x_g = *x_g * r;
146 Self::deterministic_encrypt(msg, &r_g, &r_x_g, 0)
147 }
148
149 fn decrypt(&self, sk: &G::ScalarType) -> Vec<u8> {
150 let partial_key = self.ephemeral_key * sk;
151 self.decrypt_from_partial_decryption(&partial_key)
152 }
153
154 pub fn decrypt_from_partial_decryption(&self, partial_key: &G) -> Vec<u8> {
155 let cipher = Self::sym_encrypt(partial_key, self.hkdf_info);
156 cipher
157 .decrypt(&Self::fixed_zero_nonce(), &self.data)
158 .expect("Decrypt should never fail for CTR mode")
159 }
160
161 pub fn ephemeral_key(&self) -> &G {
162 &self.ephemeral_key
163 }
164
165 fn hkdf(ikm: &G, info: usize) -> Vec<u8> {
166 let ikm = bcs::to_bytes(ikm).expect("serialize should never fail");
167 let info = info.to_be_bytes();
168 hkdf_sha3_256(
169 &HkdfIkm::from_bytes(ikm.as_slice()).expect("hkdf_sha3_256 should work with any input"),
170 &[],
171 &info,
172 AES_KEY_LENGTH,
173 )
174 .expect("hkdf_sha3_256 should never fail for an AES_KEY_LENGTH long output")
175 }
176
177 fn fixed_zero_nonce() -> InitializationVector<U16> {
178 InitializationVector::<U16>::from_bytes(&[0u8; 16])
179 .expect("U16 could always be set from a 16 bytes array of zeros")
180 }
181}
182
183impl<G: GroupElement + Serialize> MultiRecipientEncryption<G>
184where
185 <G as GroupElement>::ScalarType: FiatShamirChallenge,
186{
187 pub fn encrypt<R: AllowedRng>(
188 pk_and_msgs: &[(PublicKey<G>, Vec<u8>)],
189 random_oracle: &RandomOracle,
190 rng: &mut R,
191 ) -> MultiRecipientEncryption<G> {
192 let r = G::ScalarType::rand(rng);
193 let r_g = G::generator() * r;
194 let encs = pk_and_msgs
195 .iter()
196 .enumerate()
197 .map(|(info, (pk, msg))| {
198 let r_x_g = pk.0 * r;
199 Encryption::<G>::deterministic_encrypt(msg, &r_g, &r_x_g, info).data
200 })
201 .collect::<Vec<_>>();
202 let encs_bytes = bcs::to_bytes(&encs).expect("serialize should never fail");
204 let nizk = DLNizk::<G>::create(&r, &r_g, &encs_bytes, random_oracle, rng);
205 Self(r_g, encs, nizk)
206 }
207
208 pub fn get_encryption(&self, i: usize) -> FastCryptoResult<Encryption<G>> {
209 let buffer = self.1.get(i).ok_or(FastCryptoError::InvalidInput)?;
210 Ok(Encryption {
211 ephemeral_key: self.0,
212 data: buffer.clone(),
213 hkdf_info: i,
214 })
215 }
216
217 pub fn len(&self) -> usize {
218 self.1.len()
219 }
220 pub fn is_empty(&self) -> bool {
221 self.1.is_empty()
222 }
223
224 pub fn verify(&self, random_oracle: &RandomOracle) -> FastCryptoResult<()> {
225 let encs_bytes = bcs::to_bytes(&self.1).expect("serialize should never fail");
226 self.2.verify(&self.0, &encs_bytes, random_oracle)?;
227 self.1
229 .iter()
230 .all(|e| !e.is_empty())
231 .then_some(())
232 .ok_or(FastCryptoError::InvalidInput)
233 }
234
235 pub fn ephemeral_key(&self) -> &G {
236 &self.0
237 }
238 pub fn proof(&self) -> &DLNizk<G> {
239 &self.2
240 }
241
242 #[cfg(test)]
243 pub fn swap_for_testing(&mut self, i: usize, j: usize) {
244 self.1.swap(i, j);
245 }
246
247 #[cfg(test)]
248 pub fn copy_for_testing(&mut self, src: usize, dst: usize) {
249 self.1[dst] = self.1[src].clone();
250 }
251}