Skip to main content

confium_crypto_zk/
accumulator.rs

1//! Homomorphic hash accumulator.
2//!
3//! An RSA-style accumulator that supports:
4//! - Add element
5//! - Prove membership (witness)
6//! - Verify membership
7//! - Remove element (with trapdoor)
8
9use num_bigint::{BigUint, RandBigInt};
10use num_traits::One;
11use rand_core::OsRng;
12use sha2::{Digest, Sha256};
13use std::collections::HashMap;
14
15/// An accumulator state.
16#[derive(Debug, Clone)]
17pub struct Accumulator {
18    /// The accumulated value.
19    pub state: BigUint,
20    /// The trapdoor (prime factorization of the modulus).
21    pub trapdoor: BigUint,
22    /// The modulus N = p * q.
23    pub modulus: BigUint,
24    /// Element → prime representation.
25    pub elements: HashMap<Vec<u8>, BigUint>,
26}
27
28impl Accumulator {
29    /// Create a new accumulator with a fresh trapdoor.
30    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), // g = 2 (generator)
36            trapdoor: (&p - &BigUint::one()) * (&q - &BigUint::one()),
37            modulus: n,
38            elements: HashMap::new(),
39        }
40    }
41
42    /// Add an element to the accumulator.
43    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    /// Generate a witness (membership proof) for an element.
51    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    /// Verify membership: witness ^ element_prime == state (mod N).
64    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    /// Remove an element (requires recomputation).
71    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    /// Number of accumulated elements.
86    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    // Ensure odd
104    num |= BigUint::one();
105    // For testing, we don't verify primality — just use the hash value
106    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        // "c" is not in the accumulator
155        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    //! Paired rejects-forgery tests for membership verification.
217
218    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        // Valid witness, wrong statement: it must not verify.
226        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}