Skip to main content

binius_circuits/sha3/
varlen.rs

1// Copyright 2026 The Binius Developers
2
3use binius_core::word::Word;
4use binius_frontend::{CircuitBuilder, Wire};
5
6use super::{
7	SHA3_256_DIGEST_WORDS, SHA3_256_RATE_BYTES, SHA3_384_DIGEST_WORDS, SHA3_384_RATE_BYTES,
8	SHA3_512_DIGEST_WORDS, SHA3_512_RATE_BYTES, SHA3_DELIMITER_BYTE,
9};
10use crate::{
11	fixed_byte_vec::ByteVec,
12	keccak::{N_WORDS_PER_STATE, permutation::keccak_f1600},
13	multiplexer::{multi_wire_multiplex, single_wire_multiplex},
14};
15
16/// Computes a FIPS 202 SHA-3 digest of a variable-length message.
17///
18/// This is the shared core for every SHA-3 variant in this module.
19///
20/// Only the rate and the digest length differ between them.
21///
22/// The message length is a runtime value bounded by a fixed maximum, so the padded words are
23/// derived with runtime multiplexers instead of computed directly.
24///
25/// # Arguments
26/// * `builder` - Circuit builder for constructing constraints.
27/// * `message` - Input message with a runtime length wire and a fixed maximum capacity.
28/// * `rate_bytes` - The sponge rate for this hash variant, in bytes.
29/// * `digest_words` - The digest length for this hash variant, in 64-bit words.
30fn sha3_varlen(
31	builder: &CircuitBuilder,
32	message: &ByteVec,
33	rate_bytes: usize,
34	digest_words: usize,
35) -> Vec<Wire> {
36	let len_bytes = message.len_bytes;
37	let data = &message.data;
38
39	let n_words_per_block = rate_bytes / 8;
40	let max_len_bytes = data.len() << 3;
41	// A message that exactly fills its blocks still needs one more block for the padding.
42	let n_blocks = (max_len_bytes + 1).div_ceil(rate_bytes);
43	let n_words = n_blocks * n_words_per_block;
44
45	// Reject any claimed length past the wire capacity before it drives any derivation below.
46	let too_long = builder.icmp_ugt(len_bytes, builder.add_constant_64(max_len_bytes as u64));
47	builder.assert_false("len_check", too_long);
48
49	let zero = builder.add_constant(Word::ZERO);
50	let msb_one = builder.add_constant(Word::MSB_ONE);
51
52	// Split the length into a word index and a byte offset within that word.
53	//
54	// The word at that index is where the padding's domain suffix belongs.
55	let w_bd = builder.shr(len_bytes, 3);
56	let len_mod_8 = builder.band(len_bytes, builder.add_constant_64(7));
57
58	// Find which block holds the padding, by scanning every block's byte range.
59	//
60	// The rate is not a power of two, so this cannot be found with a shift.
61	let mut end_block_index = zero;
62	for block_no in 0..n_blocks {
63		let block_start = builder.add_constant_64((block_no * rate_bytes) as u64);
64		let block_end = builder.add_constant_64(((block_no + 1) * rate_bytes) as u64);
65		let gte_start = builder.icmp_ule(block_start, len_bytes);
66		let lt_end = builder.icmp_ult(len_bytes, block_end);
67		let is_final_block = builder.band(gte_start, lt_end);
68		end_block_index = builder.select(
69			is_final_block,
70			builder.add_constant_64(block_no as u64),
71			end_block_index,
72		);
73	}
74
75	// Build every possible boundary word, one per byte offset the length could land on.
76	//
77	// Each candidate keeps the low message bytes up to that offset and places the domain
78	// suffix right after them.
79	//
80	// At offset zero the message contributes nothing, so the candidate is the suffix alone.
81	let boundary_message_word = single_wire_multiplex(builder, data, w_bd);
82	let candidates: Vec<Wire> = (0..8)
83		.map(|i| {
84			let mask = builder.add_constant_64(0x00FFFFFFFFFFFFFF >> ((7 - i) << 3));
85			let delimiter = builder.add_constant_64(SHA3_DELIMITER_BYTE << (i << 3));
86			let message_low = builder.band(boundary_message_word, mask);
87			builder.bxor(message_low, delimiter)
88		})
89		.collect();
90	// Select the one candidate that matches the actual byte offset.
91	let boundary_word = single_wire_multiplex(builder, &candidates, len_mod_8);
92
93	// Derive every padded word, one at a time, classified by its position relative to the
94	// boundary word and to the block holding the padding.
95	let padded_message: Vec<Wire> = (0..n_words)
96		.map(|word_index| {
97			let block_index = word_index / n_words_per_block;
98			let column_index = word_index % n_words_per_block;
99			let word_idx_wire = builder.add_constant_64(word_index as u64);
100
101			let is_message_word = builder.icmp_ult(word_idx_wire, w_bd);
102			let is_boundary_word = builder.icmp_eq(word_idx_wire, w_bd);
103			let is_end_block =
104				builder.icmp_eq(builder.add_constant_64(block_index as u64), end_block_index);
105
106			// A word strictly before the boundary is plain message data.
107			//
108			// This index is only ever selected when it is within the message capacity, so the
109			// zero fallback below it is never actually read.
110			let msg_word = if word_index < data.len() {
111				data[word_index]
112			} else {
113				zero
114			};
115
116			// The last word of the last block carries the padding rule's closing bit.
117			//
118			// It combines with the boundary word when the boundary falls there too, and
119			// otherwise stands alone as a plain padding word.
120			let delimiter = if column_index == n_words_per_block - 1 {
121				builder.select(is_end_block, msb_one, zero)
122			} else {
123				zero
124			};
125			let boundary_val = if column_index == n_words_per_block - 1 {
126				builder.bxor(boundary_word, delimiter)
127			} else {
128				boundary_word
129			};
130
131			let boundary_or_padding = builder.select(is_boundary_word, boundary_val, delimiter);
132			builder.select(is_message_word, msg_word, boundary_or_padding)
133		})
134		.collect();
135
136	// Absorb one padded block at a time.
137	//
138	// XOR it into the front of the state, then permute the whole state.
139	//
140	// Every intermediate state is kept, since the block holding the padding is only known at
141	// runtime.
142	let mut states: Vec<[Wire; N_WORDS_PER_STATE]> = Vec::with_capacity(n_blocks + 1);
143	states.push([zero; N_WORDS_PER_STATE]);
144	for block_no in 0..n_blocks {
145		let mut state = states[block_no];
146		for i in 0..n_words_per_block {
147			state[i] = builder.bxor(state[i], padded_message[block_no * n_words_per_block + i]);
148		}
149		keccak_f1600(builder, &mut state);
150		states.push(state);
151	}
152
153	// Select the digest out of the one state that followed the block holding the padding.
154	let inputs: Vec<&[Wire]> = states[1..].iter().map(|s| &s[..]).collect();
155	let digest_vec = multi_wire_multiplex(builder, &inputs, end_block_index);
156	digest_vec[..digest_words].to_vec()
157}
158
159/// Computes the SHA3-256 digest of a variable-length message.
160///
161/// # Arguments
162/// * `builder` - Circuit builder for constructing constraints.
163/// * `message` - Input message with a runtime length wire and a fixed maximum capacity.
164///
165/// # Returns
166/// The digest as 4 little-endian 64-bit words.
167pub fn sha3_256_varlen(
168	builder: &CircuitBuilder,
169	message: &ByteVec,
170) -> [Wire; SHA3_256_DIGEST_WORDS] {
171	sha3_varlen(builder, message, SHA3_256_RATE_BYTES, SHA3_256_DIGEST_WORDS)
172		.try_into()
173		.unwrap()
174}
175
176/// Computes the SHA3-384 digest of a variable-length message.
177///
178/// # Arguments
179/// * `builder` - Circuit builder for constructing constraints.
180/// * `message` - Input message with a runtime length wire and a fixed maximum capacity.
181///
182/// # Returns
183/// The digest as 6 little-endian 64-bit words.
184pub fn sha3_384_varlen(
185	builder: &CircuitBuilder,
186	message: &ByteVec,
187) -> [Wire; SHA3_384_DIGEST_WORDS] {
188	sha3_varlen(builder, message, SHA3_384_RATE_BYTES, SHA3_384_DIGEST_WORDS)
189		.try_into()
190		.unwrap()
191}
192
193/// Computes the SHA3-512 digest of a variable-length message.
194///
195/// # Arguments
196/// * `builder` - Circuit builder for constructing constraints.
197/// * `message` - Input message with a runtime length wire and a fixed maximum capacity.
198///
199/// # Returns
200/// The digest as 8 little-endian 64-bit words.
201pub fn sha3_512_varlen(
202	builder: &CircuitBuilder,
203	message: &ByteVec,
204) -> [Wire; SHA3_512_DIGEST_WORDS] {
205	sha3_varlen(builder, message, SHA3_512_RATE_BYTES, SHA3_512_DIGEST_WORDS)
206		.try_into()
207		.unwrap()
208}
209
210#[cfg(test)]
211mod tests {
212	use binius_core::Word;
213	use binius_frontend::{CircuitBuilder, Wire};
214	use rand::prelude::*;
215	use rstest::rstest;
216	use sha3::Digest;
217
218	use super::*;
219
220	// Builds a circuit with the given capacity, runs one hash function on a message that fits
221	// within it, and checks the computed digest against a reference implementation.
222	fn test_varlen<const DIGEST_WORDS: usize>(
223		message: &[u8],
224		max_message_len_bytes: usize,
225		hash_fn: impl FnOnce(&CircuitBuilder, &ByteVec) -> [Wire; DIGEST_WORDS],
226		reference: impl FnOnce(&[u8]) -> Vec<u8>,
227	) {
228		assert!(message.len() <= max_message_len_bytes);
229
230		let expected_digest = reference(message);
231		assert_eq!(expected_digest.len(), DIGEST_WORDS * 8);
232
233		let b = CircuitBuilder::new();
234		let max_len_words = max_message_len_bytes.div_ceil(8);
235		let input = ByteVec::new_inout(&b, max_len_words);
236		let expected_digest_wires: [Wire; DIGEST_WORDS] = std::array::from_fn(|_| b.add_witness());
237
238		// Constrain the circuit's digest to equal the expected one, wire by wire.
239		let computed_digest = hash_fn(&b, &input);
240		for i in 0..DIGEST_WORDS {
241			b.assert_eq(format!("digest[{i}]"), computed_digest[i], expected_digest_wires[i]);
242		}
243
244		let circuit = b.build();
245		let cs = circuit.constraint_system();
246		let mut witness = circuit.new_witness_filler();
247
248		// Populate the message below its capacity and record its true length.
249		input.populate_data(&mut witness, message);
250		input.populate_len_bytes(&mut witness, message.len());
251		// Populate the expected digest wires.
252		for (i, bytes) in expected_digest.chunks(8).enumerate() {
253			witness[expected_digest_wires[i]] = Word(u64::from_le_bytes(bytes.try_into().unwrap()));
254		}
255
256		// Evaluating the circuit fills in every internal wire from the message and digest inputs.
257		circuit
258			.populate_wire_witness(&mut witness)
259			.expect("Circuit should accept valid witness");
260		// Checking the constraint system confirms the computed digest equals the reference
261		// digest, not just that the witness filled in without panicking.
262		cs.verify(&witness.into_value_vec())
263			.expect("All constraints should be satisfied");
264	}
265
266	// Deterministic random message of the given length, seeded from both lengths so every
267	// combination of message length and capacity gets an independent, reproducible message.
268	fn random_message(message_len_bytes: usize, max_message_len_bytes: usize) -> Vec<u8> {
269		let seed = ((message_len_bytes as u64) << 32) | (max_message_len_bytes as u64);
270		let mut rng = StdRng::seed_from_u64(seed);
271		let mut message = vec![0u8; message_len_bytes];
272		rng.fill_bytes(&mut message);
273		message
274	}
275
276	#[rstest]
277	#[case(0, 100)] // Empty message
278	#[case(1, 100)] // Single byte, well below capacity
279	#[case(1, 144)] // Single byte, capacity spans two blocks
280	#[case(135, 136)] // One byte before the block boundary
281	#[case(136, 136)] // Exactly one block
282	#[case(137, 272)] // Crosses the block boundary
283	#[case(271, 272)] // One byte before two blocks
284	#[case(272, 272)] // Exactly two blocks
285	fn test_sha3_256_varlen(
286		#[case] message_len_bytes: usize,
287		#[case] max_message_len_bytes: usize,
288	) {
289		let message = random_message(message_len_bytes, max_message_len_bytes);
290		test_varlen(&message, max_message_len_bytes, sha3_256_varlen, |m| {
291			sha3::Sha3_256::digest(m).to_vec()
292		});
293	}
294
295	#[rstest]
296	#[case(0, 100)] // Empty message
297	#[case(1, 100)] // Single byte, well below capacity
298	#[case(103, 104)] // One byte before the block boundary
299	#[case(104, 104)] // Exactly one block
300	#[case(105, 208)] // Crosses the block boundary
301	fn test_sha3_384_varlen(
302		#[case] message_len_bytes: usize,
303		#[case] max_message_len_bytes: usize,
304	) {
305		let message = random_message(message_len_bytes, max_message_len_bytes);
306		test_varlen(&message, max_message_len_bytes, sha3_384_varlen, |m| {
307			sha3::Sha3_384::digest(m).to_vec()
308		});
309	}
310
311	#[rstest]
312	#[case(0, 100)] // Empty message
313	#[case(1, 100)] // Single byte, well below capacity
314	#[case(71, 72)] // One byte before the block boundary
315	#[case(72, 72)] // Exactly one block
316	#[case(73, 144)] // Crosses the block boundary
317	fn test_sha3_512_varlen(
318		#[case] message_len_bytes: usize,
319		#[case] max_message_len_bytes: usize,
320	) {
321		let message = random_message(message_len_bytes, max_message_len_bytes);
322		test_varlen(&message, max_message_len_bytes, sha3_512_varlen, |m| {
323			sha3::Sha3_512::digest(m).to_vec()
324		});
325	}
326}