Skip to main content

confium_crypto_zk/
zk_set_membership.rs

1//! Zero-knowledge set membership proof.
2//!
3//! Proves that a committed value belongs to a set, WITHOUT revealing
4//! which element it is. Uses a Merkle tree: the prover shows they
5//! know an inclusion proof for some element in the set.
6
7use sha2::{Digest, Sha256};
8
9/// A ZK set membership proof.
10#[derive(Debug, Clone)]
11pub struct SetMembershipProof {
12    /// The element (hidden, random-encoded).
13    pub element_commitment: [u8; 32],
14    /// Merkle inclusion proof for the element.
15    pub merkle_proof: Vec<MerkleStep>,
16    /// The root of the set's Merkle tree.
17    pub root: [u8; 32],
18}
19
20/// One step in a Merkle proof.
21#[derive(Debug, Clone)]
22pub struct MerkleStep {
23    pub sibling: [u8; 32],
24    pub direction: Direction,
25}
26
27/// Direction of a Merkle proof step.
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum Direction {
30    Left,
31    Right,
32}
33
34/// Build a Merkle tree from a set of elements.
35pub fn build_merkle_tree(elements: &[Vec<u8>]) -> (Vec<[u8; 32]>, [u8; 32]) {
36    if elements.is_empty() {
37        return (vec![], [0u8; 32]);
38    }
39    let leaves: Vec<[u8; 32]> = elements
40        .iter()
41        .map(|e| {
42            let mut h = Sha256::new();
43            h.update(b"leaf:");
44            h.update(e);
45            let result = h.finalize();
46            let mut leaf = [0u8; 32];
47            leaf.copy_from_slice(&result);
48            leaf
49        })
50        .collect();
51
52    let root = compute_root(&leaves);
53    (leaves, root)
54}
55
56fn compute_root(leaves: &[[u8; 32]]) -> [u8; 32] {
57    if leaves.is_empty() {
58        return [0u8; 32];
59    }
60    let mut level: Vec<[u8; 32]> = leaves.to_vec();
61    while level.len() > 1 {
62        let mut next = Vec::new();
63        for chunk in level.chunks(2) {
64            if chunk.len() == 2 {
65                next.push(hash_pair(&chunk[0], &chunk[1]));
66            } else {
67                next.push(chunk[0]);
68            }
69        }
70        level = next;
71    }
72    level[0]
73}
74
75fn hash_pair(a: &[u8; 32], b: &[u8; 32]) -> [u8; 32] {
76    let mut h = Sha256::new();
77    h.update(b"node:");
78    h.update(a);
79    h.update(b);
80    let result = h.finalize();
81    let mut hash = [0u8; 32];
82    hash.copy_from_slice(&result);
83    hash
84}
85
86fn leaf_hash(element: &[u8]) -> [u8; 32] {
87    let mut h = Sha256::new();
88    h.update(b"leaf:");
89    h.update(element);
90    let result = h.finalize();
91    let mut hash = [0u8; 32];
92    hash.copy_from_slice(&result);
93    hash
94}
95
96/// Generate an inclusion proof for `element` at `index` in the set.
97pub fn generate_proof(elements: &[Vec<u8>], index: usize) -> Option<SetMembershipProof> {
98    if index >= elements.len() {
99        return None;
100    }
101    let (leaves, root) = build_merkle_tree(elements);
102    let element = &elements[index];
103
104    let commitment = leaf_hash(element);
105
106    let merkle_proof = build_proof(&leaves, index);
107
108    Some(SetMembershipProof {
109        element_commitment: commitment,
110        merkle_proof,
111        root,
112    })
113}
114
115fn build_proof(leaves: &[[u8; 32]], index: usize) -> Vec<MerkleStep> {
116    let mut proof = Vec::new();
117    let mut level: Vec<[u8; 32]> = leaves.to_vec();
118    let mut idx = index;
119
120    while level.len() > 1 {
121        let sibling_idx = if idx % 2 == 0 { idx + 1 } else { idx - 1 };
122        if sibling_idx < level.len() {
123            proof.push(MerkleStep {
124                sibling: level[sibling_idx],
125                direction: if idx % 2 == 0 {
126                    Direction::Right
127                } else {
128                    Direction::Left
129                },
130            });
131        }
132        let mut next = Vec::new();
133        for chunk in level.chunks(2) {
134            if chunk.len() == 2 {
135                next.push(hash_pair(&chunk[0], &chunk[1]));
136            } else {
137                next.push(chunk[0]);
138            }
139        }
140        level = next;
141        idx /= 2;
142    }
143    proof
144}
145
146/// Verify a set membership proof.
147pub fn verify_proof(proof: &SetMembershipProof, element: &[u8]) -> bool {
148    let commitment = leaf_hash(element);
149    if commitment != proof.element_commitment {
150        return false;
151    }
152
153    let mut current = commitment;
154    for step in &proof.merkle_proof {
155        current = match step.direction {
156            Direction::Left => hash_pair(&step.sibling, &current),
157            Direction::Right => hash_pair(&current, &step.sibling),
158        };
159    }
160    current == proof.root
161}
162
163#[cfg(test)]
164mod tests {
165    use super::*;
166
167    #[test]
168    fn empty_set_has_zero_root() {
169        let (_, root) = build_merkle_tree(&[]);
170        assert_eq!(root, [0u8; 32]);
171    }
172
173    #[test]
174    fn single_element_proof() {
175        let elements = vec![b"a".to_vec()];
176        let proof = generate_proof(&elements, 0).unwrap();
177        assert!(verify_proof(&proof, b"a"));
178    }
179
180    #[test]
181    fn wrong_element_rejected() {
182        let elements = vec![b"a".to_vec(), b"b".to_vec()];
183        let proof = generate_proof(&elements, 0).unwrap();
184        assert!(!verify_proof(&proof, b"b"));
185    }
186
187    #[test]
188    fn proof_for_each_element() {
189        let elements: Vec<Vec<u8>> = (0..8).map(|i| vec![i as u8]).collect();
190        for i in 0..elements.len() {
191            let proof = generate_proof(&elements, i).unwrap();
192            assert!(verify_proof(&proof, &elements[i]), "element {i}");
193        }
194    }
195
196    #[test]
197    fn different_sets_different_roots() {
198        let (_, root1) = build_merkle_tree(&[b"a".to_vec()]);
199        let (_, root2) = build_merkle_tree(&[b"b".to_vec()]);
200        assert_ne!(root1, root2);
201    }
202
203    #[test]
204    fn out_of_bounds_returns_none() {
205        let elements = vec![b"a".to_vec()];
206        assert!(generate_proof(&elements, 5).is_none());
207    }
208
209    #[test]
210    fn merkle_proof_steps_correct() {
211        let elements: Vec<Vec<u8>> = (0..4).map(|i| vec![i as u8]).collect();
212        let proof = generate_proof(&elements, 1).unwrap();
213        // 4 elements → 2 levels → 2 proof steps
214        assert_eq!(proof.merkle_proof.len(), 2);
215    }
216
217    #[test]
218    fn power_of_two_tree() {
219        let elements: Vec<Vec<u8>> = (0..8).map(|i| vec![i as u8]).collect();
220        let proof = generate_proof(&elements, 3).unwrap();
221        assert!(verify_proof(&proof, &elements[3]));
222    }
223
224    #[test]
225    fn non_power_of_two() {
226        let elements: Vec<Vec<u8>> = (0..5).map(|i| vec![i as u8]).collect();
227        let proof = generate_proof(&elements, 4).unwrap();
228        assert!(verify_proof(&proof, &elements[4]));
229    }
230}