1use crate::{
2 AllWire, ArithmeticWireLabel, BinaryWireLabel, WireLabel, WireMod2,
3 util::{output_tweak, tweak2},
4 wire::hash_wires,
5};
6use fancy_traits::{
7 Fancy, FancyArithmetic, FancyBinary, FancyBinaryConstant, FancyConstant, FancyEncode,
8 FancyOutput, HasModulus, is_binary,
9};
10use rand::{CryptoRng, RngExt};
11#[cfg(feature = "serde")]
12use serde::de::DeserializeOwned;
13use std::collections::HashMap;
14use swanky_channel::Channel;
15use swanky_field_binary::F2;
16use vectoreyes::U8x16;
17
18pub struct Garbler<RNG, Wire> {
20 zero: Wire,
22 delta_mod_2: Wire,
23 deltas: HashMap<u16, Wire>,
25 current_output: usize,
26 current_gate: usize,
27 rng: RNG,
28}
29
30#[cfg(feature = "serde")]
31impl<RNG: CryptoRng, Wire: WireLabel + DeserializeOwned> Garbler<RNG, Wire> {
32 pub fn load_deltas(&mut self, filename: &str) -> Result<(), Box<dyn std::error::Error>> {
34 let f = std::fs::File::open(filename)?;
35 let reader = std::io::BufReader::new(f);
36 let deltas: HashMap<u16, Wire> = serde_json::from_reader(reader)?;
37 self.deltas.extend(deltas);
38 Ok(())
39 }
40}
41
42impl<RNG: CryptoRng, Wire: WireLabel> Garbler<RNG, Wire> {
43 pub fn new(mut rng: RNG) -> Self {
45 let delta = Wire::rand_delta(&mut rng, 2);
46 let one = Wire::from_repr(U8x16::from(1u128), 2);
49 let zero = delta.clone() + one;
50 Garbler {
51 zero,
52 delta_mod_2: delta,
53 deltas: HashMap::new(),
54 current_gate: 0,
55 current_output: 0,
56 rng,
57 }
58 }
59
60 fn current_gate(&mut self) -> usize {
62 let current = self.current_gate;
63 self.current_gate += 1;
64 current
65 }
66
67 pub fn delta(&mut self, q: u16) -> Wire {
70 if q == 2 {
71 self.delta_mod_2.clone()
72 } else if let Some(delta) = self.deltas.get(&q) {
73 delta.clone()
74 } else {
75 let w = Wire::rand_delta(&mut self.rng, q);
76 self.deltas.insert(q, w.clone());
77 w
78 }
79 }
80
81 fn current_output(&mut self) -> usize {
83 let current = self.current_output;
84 self.current_output += 1;
85 current
86 }
87
88 pub fn get_deltas(mut self) -> HashMap<u16, Wire> {
90 self.deltas.insert(2, self.delta_mod_2);
92 self.deltas
93 }
94
95 pub fn encode_zero(&mut self, modulus: u16) -> Wire {
97 Wire::rand(&mut self.rng, modulus)
98 }
99}
100
101impl<RNG: CryptoRng, W: BinaryWireLabel> FancyBinary for Garbler<RNG, W> {
102 fn and(
103 &mut self,
104 A: &Self::Item,
105 B: &Self::Item,
106 channel: &mut Channel,
107 ) -> swanky_error::Result<Self::Item> {
108 let delta = self.delta(2);
109 let gate_num = self.current_gate();
110 let (gate0, gate1, C) = W::garble_and_gate(gate_num, A, B, &delta);
111 channel.write(&gate0)?;
112 channel.write(&gate1)?;
113 Ok(C)
114 }
115
116 fn xor(&mut self, x: &Self::Item, y: &Self::Item) -> Self::Item {
117 *x + *y
118 }
119
120 fn negate(&mut self, x: &Self::Item) -> Self::Item {
125 self.zero + *x
126 }
127}
128
129impl<RNG: CryptoRng> FancyBinary for Garbler<RNG, AllWire> {
130 fn negate(&mut self, x: &Self::Item) -> Self::Item {
135 is_binary!(x);
136
137 let zero = self.zero.clone();
138 self.xor(&zero, x)
139 }
140
141 fn xor(&mut self, x: &Self::Item, y: &Self::Item) -> Self::Item {
143 is_binary!(x);
144 is_binary!(y);
145
146 self.add(x, y)
147 }
148
149 fn and(
151 &mut self,
152 x: &Self::Item,
153 y: &Self::Item,
154 channel: &mut Channel,
155 ) -> swanky_error::Result<Self::Item> {
156 if let (AllWire::Mod2(A), AllWire::Mod2(B), AllWire::Mod2(ref delta)) =
157 (x, y, self.delta(2))
158 {
159 let gate_num = self.current_gate();
160 let (gate0, gate1, C) = WireMod2::garble_and_gate(gate_num, A, B, delta);
161 channel.write(&gate0)?;
162 channel.write(&gate1)?;
163 return Ok(AllWire::Mod2(C));
164 }
165 is_binary!(x);
167 is_binary!(y);
168
169 unreachable!()
171 }
172}
173
174impl<RNG: CryptoRng, Wire: WireLabel + ArithmeticWireLabel> FancyArithmetic for Garbler<RNG, Wire> {
175 fn add(&mut self, x: &Wire, y: &Wire) -> Wire {
176 assert_eq!(x.modulus(), y.modulus());
177 x.clone() + y.clone()
178 }
179
180 fn sub(&mut self, x: &Wire, y: &Wire) -> Wire {
181 assert_eq!(x.modulus(), y.modulus());
182 x.clone() - y.clone()
183 }
184
185 fn cmul(&mut self, x: &Wire, c: u16) -> Wire {
186 x.clone() * c
187 }
188
189 fn mul(&mut self, A: &Wire, B: &Wire, channel: &mut Channel) -> swanky_error::Result<Wire> {
190 if A.modulus() < B.modulus() {
191 return self.mul(B, A, channel);
192 }
193
194 let q = A.modulus();
195 let qb = B.modulus();
196 let gate_num = self.current_gate();
197
198 let D = self.delta(q);
199 let Db = self.delta(qb);
200
201 let r;
202 let mut gate = vec![Default::default(); q as usize + qb as usize - 2];
203
204 if q != qb {
206 assert!(
208 qb <= 8,
209 "`B.modulus()` with asymmetric moduli is capped at 8"
210 );
211
212 r = self.rng.random::<u16>() % q;
213 let t = tweak2(gate_num as u64, 1);
214
215 let mut minitable = vec![u128::default(); qb as usize];
216 let mut B_ = B.clone();
217 for b in 0..qb {
218 if b > 0 {
219 B_ += Db.clone();
220 }
221 let new_color = ((r + b) % q) as u128;
222 let ct = (u128::from(B_.hash(t)) & 0xFFFF) ^ new_color;
223 minitable[B_.color() as usize] = ct;
224 }
225
226 let mut packed = 0;
227 for (i, item) in minitable.iter().enumerate().take(qb as usize) {
228 packed += item << (16 * i);
229 }
230 gate.push(packed.into());
231 } else {
232 r = B.color(); }
234
235 let g = tweak2(gate_num as u64, 0);
236
237 let alpha = (q - A.color()) % q; let X1 = A.clone() + D.clone() * alpha;
240
241 let beta = (qb - B.color()) % qb;
243 let Y1 = B.clone() + Db.clone() * beta;
244
245 let [hashX, hashY] = hash_wires([&X1, &Y1], g);
246
247 let X = Wire::hash_to_mod(hashX, q) + D.clone() * (alpha * r % q);
248 let Y = Wire::hash_to_mod(hashY, q) + A.clone() * ((beta + r) % q);
249
250 let mut precomp = Vec::with_capacity(q as usize);
251 let mut X_ = X.clone();
254 precomp.push(X_.to_repr());
255 for _ in 1..q {
256 X_ += D.clone();
257 precomp.push(X_.to_repr());
258 }
259
260 let mut A_ = A.clone();
264 for a in 0..q {
265 if a > 0 {
266 A_ += D.clone();
267 }
268 if A_.color() != 0 {
271 gate[A_.color() as usize - 1] =
272 A_.hash(g) ^ precomp[((q - (a * r % q)) % q) as usize];
273 }
274 }
275 precomp.clear();
276
277 let mut Y_ = Y.clone();
280 precomp.push(Y_.to_repr());
281 for _ in 1..q {
282 Y_ += A.clone();
283 precomp.push(Y_.to_repr());
284 }
285
286 let mut B_ = B.clone();
288 for b in 0..qb {
289 if b > 0 {
290 B_ += Db.clone();
291 }
292 if B_.color() != 0 {
295 gate[q as usize - 1 + B_.color() as usize - 1] =
296 B_.hash(g) ^ precomp[((q - ((b + r) % q)) % q) as usize];
297 }
298 }
299
300 for block in gate.iter() {
301 channel.write(block)?;
302 }
303 Ok(X + Y)
304 }
305}
306
307impl<RNG: CryptoRng, Wire: WireLabel> Fancy for Garbler<RNG, Wire> {
308 type Item = Wire;
309}
310
311impl<RNG: CryptoRng, Wire: WireLabel> FancyConstant for Garbler<RNG, Wire> {
312 fn constant(&mut self, x: u16, q: u16, channel: &mut Channel) -> swanky_error::Result<Wire> {
313 let (zero, wire) = Wire::constant(x, q, &self.delta(q), &mut self.rng);
314 channel.write(&wire.to_repr())?;
315 Ok(zero)
316 }
317}
318
319impl<RNG: CryptoRng, Wire: WireLabel> FancyBinaryConstant for Garbler<RNG, Wire> {
320 fn constant(&mut self, x: F2) -> Self::Item {
321 if x.into() {
322 self.zero.clone()
325 } else {
326 Default::default()
328 }
329 }
330}
331
332impl<RNG: CryptoRng, Wire: WireLabel> FancyEncode for Garbler<RNG, Wire> {
333 fn encode_many(
334 &mut self,
335 values: &[u16],
336 moduli: &[u16],
337 channel: &mut Channel,
338 ) -> swanky_error::Result<Vec<Self::Item>> {
339 assert_eq!(values.len(), moduli.len());
340
341 let mut zeros = Vec::with_capacity(values.len());
342 for (x, q) in values.iter().zip(moduli.iter()) {
343 let delta = self.delta(*q);
344 let zero = self.encode_zero(*q);
345 let encoded = zero.clone() + delta * *x;
346 channel.write(&encoded.to_repr())?;
347 zeros.push(zero);
348 }
349 Ok(zeros)
350 }
351
352 fn receive_many(
353 &mut self,
354 _moduli: &[u16],
355 _: &mut Channel,
356 ) -> swanky_error::Result<Vec<Self::Item>> {
357 unimplemented!("Garbler cannot receive values")
358 }
359}
360
361impl<RNG: CryptoRng, Wire: WireLabel> FancyOutput for Garbler<RNG, Wire> {
362 fn output(&mut self, X: &Wire, channel: &mut Channel) -> swanky_error::Result<Option<u16>> {
363 let q = X.modulus();
364 let i = self.current_output();
365 let D = self.delta(q);
366 for k in 0..q {
367 let block = (X.clone() + D.clone() * k).hash(output_tweak(i, k));
368 channel.write(&block)?;
369 }
370 Ok(None)
371 }
372}