confium_crypto_zk/
accumulator.rs1use num_bigint::{BigUint, RandBigInt};
10use num_traits::One;
11use rand_core::OsRng;
12use sha2::{Digest, Sha256};
13use std::collections::HashMap;
14
15#[derive(Debug, Clone)]
17pub struct Accumulator {
18 pub state: BigUint,
20 pub trapdoor: BigUint,
22 pub modulus: BigUint,
24 pub elements: HashMap<Vec<u8>, BigUint>,
26}
27
28impl Accumulator {
29 pub fn new() -> Self {
31 let p = generate_prime(128);
32 let q = generate_prime(128);
33 let n = &p * &q;
34 Self {
35 state: BigUint::from(2u32), trapdoor: (&p - &BigUint::one()) * (&q - &BigUint::one()),
37 modulus: n,
38 elements: HashMap::new(),
39 }
40 }
41
42 pub fn add(&mut self, element: &[u8]) -> BigUint {
44 let prime = hash_to_prime(element);
45 self.state = self.state.modpow(&prime, &self.modulus);
46 self.elements.insert(element.to_vec(), prime.clone());
47 prime
48 }
49
50 pub fn witness(&self, element: &[u8]) -> Option<BigUint> {
52 let target_prime = self.elements.get(element)?;
53 let mut product = BigUint::one();
54 for prime in self.elements.values() {
55 if prime != target_prime {
56 product *= prime;
57 }
58 }
59 let g = BigUint::from(2u32);
60 Some(g.modpow(&product, &self.modulus))
61 }
62
63 pub fn verify(&self, witness: &BigUint, element: &[u8]) -> bool {
65 let prime = hash_to_prime(element);
66 let expected = witness.modpow(&prime, &self.modulus);
67 expected == self.state
68 }
69
70 pub fn remove(&mut self, element: &[u8]) -> bool {
72 if self.elements.remove(element).is_some() {
73 let g = BigUint::from(2u32);
74 let mut state = g;
75 for prime in self.elements.values() {
76 state = state.modpow(prime, &self.modulus);
77 }
78 self.state = state;
79 true
80 } else {
81 false
82 }
83 }
84
85 pub fn count(&self) -> usize {
87 self.elements.len()
88 }
89}
90
91impl Default for Accumulator {
92 fn default() -> Self {
93 Self::new()
94 }
95}
96
97fn hash_to_prime(element: &[u8]) -> BigUint {
98 let mut h = Sha256::new();
99 h.update(b"acc-prime:");
100 h.update(element);
101 let result = h.finalize();
102 let mut num = BigUint::from_bytes_be(&result);
103 num |= BigUint::one();
105 num
107}
108
109fn generate_prime(bits: u32) -> BigUint {
110 let mut rng = OsRng;
111 loop {
112 let candidate = rng.gen_biguint(bits as u64);
113 if candidate > BigUint::from(3u32) {
114 return candidate | BigUint::one();
115 }
116 }
117}
118
119#[cfg(test)]
120mod tests {
121 use super::*;
122
123 #[test]
124 fn new_accumulator_empty() {
125 let acc = Accumulator::new();
126 assert_eq!(acc.count(), 0);
127 assert!(acc.state > BigUint::one());
128 }
129
130 #[test]
131 fn add_changes_state() {
132 let mut acc = Accumulator::new();
133 let initial = acc.state.clone();
134 acc.add(b"element");
135 assert_ne!(acc.state, initial);
136 assert_eq!(acc.count(), 1);
137 }
138
139 #[test]
140 fn witness_verifies_membership() {
141 let mut acc = Accumulator::new();
142 acc.add(b"a");
143 acc.add(b"b");
144 acc.add(b"c");
145 let witness = acc.witness(b"b").unwrap();
146 assert!(acc.verify(&witness, b"b"));
147 }
148
149 #[test]
150 fn non_member_not_verified() {
151 let mut acc = Accumulator::new();
152 acc.add(b"a");
153 acc.add(b"b");
154 let fake_witness = BigUint::from(2u32);
156 assert!(!acc.verify(&fake_witness, b"c"));
157 }
158
159 #[test]
160 fn witness_for_each_element() {
161 let mut acc = Accumulator::new();
162 let elements: Vec<Vec<u8>> = (0..5).map(|i| vec![i as u8]).collect();
163 for e in &elements {
164 acc.add(e);
165 }
166 for e in &elements {
167 let w = acc.witness(e).unwrap();
168 assert!(acc.verify(&w, e), "element {:?}", e);
169 }
170 }
171
172 #[test]
173 fn remove_decrements_count() {
174 let mut acc = Accumulator::new();
175 acc.add(b"a");
176 acc.add(b"b");
177 assert_eq!(acc.count(), 2);
178 assert!(acc.remove(b"a"));
179 assert_eq!(acc.count(), 1);
180 }
181
182 #[test]
183 fn remove_unknown_returns_false() {
184 let mut acc = Accumulator::new();
185 acc.add(b"a");
186 assert!(!acc.remove(b"z"));
187 }
188
189 #[test]
190 fn witness_unknown_returns_none() {
191 let acc = Accumulator::new();
192 assert!(acc.witness(b"unknown").is_none());
193 }
194
195 #[test]
196 fn re_add_works_after_remove() {
197 let mut acc = Accumulator::new();
198 acc.add(b"a");
199 acc.add(b"b");
200 acc.remove(b"a");
201 acc.add(b"a");
202 let w = acc.witness(b"a").unwrap();
203 assert!(acc.verify(&w, b"a"));
204 }
205
206 #[test]
207 fn same_element_same_prime() {
208 let p1 = hash_to_prime(b"test");
209 let p2 = hash_to_prime(b"test");
210 assert_eq!(p1, p2);
211 }
212}
213
214#[cfg(test)]
215mod adversarial_tests {
216 use super::*;
219
220 #[test]
221 fn verify_rejects_wrong_element() {
222 let mut acc = Accumulator::new();
223 acc.add(b"member-a");
224 let witness = acc.witness(b"member-a").unwrap();
225 assert!(!acc.verify(&witness, b"not-a-member"));
227 }
228
229 #[test]
230 fn verify_rejects_tampered_witness() {
231 let mut acc = Accumulator::new();
232 acc.add(b"member-a");
233 let mut witness = acc.witness(b"member-a").unwrap();
234 witness += BigUint::one();
235 assert!(!acc.verify(&witness, b"member-a"));
236 }
237}