confium_crypto_zk/
zk_set_membership.rs1use sha2::{Digest, Sha256};
8
9#[derive(Debug, Clone)]
11pub struct SetMembershipProof {
12 pub element_commitment: [u8; 32],
14 pub merkle_proof: Vec<MerkleStep>,
16 pub root: [u8; 32],
18}
19
20#[derive(Debug, Clone)]
22pub struct MerkleStep {
23 pub sibling: [u8; 32],
24 pub direction: Direction,
25}
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum Direction {
30 Left,
31 Right,
32}
33
34pub 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
96pub 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
146pub 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, ¤t),
157 Direction::Right => hash_pair(¤t, &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 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}