1use alloc::vec::Vec;
9use core::marker::PhantomData;
10
11use miden_core::{
12 Felt, ZERO,
13 deferred::{
14 DeferredContext, Digest, Node, NodeType, Payload, Precompile, PrecompileError, Tag,
15 precompile_id,
16 },
17};
18
19use crate::codec::{chunks_to_bytes_exact, n_chunks};
20
21pub mod keccak256;
22
23pub trait HashFunction: Default + Send + Sync + 'static {
28 const NAME: &'static str;
30 const DIGEST_FELTS: usize;
32 fn hash(input: &[u8]) -> Vec<u8>;
34}
35
36const ASSERT_DISC: u32 = 0;
40
41#[derive(Debug, Clone, Copy, PartialEq, Eq)]
46pub struct HashAssertNode {
47 pub n_bytes: u32,
49 pub preimage_digest: Digest,
51 pub expected_digest: Digest,
53}
54
55pub struct HashPrecompile<H>(PhantomData<H>);
57
58impl<H> Default for HashPrecompile<H> {
59 fn default() -> Self {
60 Self(PhantomData)
61 }
62}
63
64impl<H: HashFunction> HashPrecompile<H> {
65 pub const ASSERT_TAG_ID: u32 = ASSERT_DISC;
67
68 pub fn id() -> Felt {
70 precompile_id(H::NAME)
71 }
72
73 pub fn assert_tag(n_bytes: u32) -> Tag {
75 Self::tag([Felt::from_u32(ASSERT_DISC), Felt::from_u32(n_bytes), ZERO])
76 }
77
78 pub fn assert_node(n_bytes: u32, preimage_digest: Digest, expected_digest: Digest) -> Node {
80 Node::join(Self::assert_tag(n_bytes), preimage_digest, expected_digest)
81 .expect("assert tag is precompile-owned")
82 }
83
84 pub fn decode_assert_tag(tag: Tag) -> Result<Option<u32>, PrecompileError> {
89 if tag.id() != Self::id() {
90 return Ok(None);
91 }
92
93 let args = tag.args();
94 let disc =
95 u32::try_from(args[0].as_canonical_u64()).map_err(|_| PrecompileError::InvalidNode)?;
96 let n_bytes =
97 u32::try_from(args[1].as_canonical_u64()).map_err(|_| PrecompileError::InvalidNode)?;
98 if disc != ASSERT_DISC || args[2] != ZERO {
99 return Err(PrecompileError::InvalidNode);
100 }
101
102 Ok(Some(n_bytes))
103 }
104
105 pub fn decode_assert_node(node: &Node) -> Result<Option<HashAssertNode>, PrecompileError> {
110 let Some(n_bytes) = Self::decode_assert_tag(node.tag())? else {
111 return Ok(None);
112 };
113 let (preimage_digest, expected_digest) = node.payload().as_join()?;
114 Ok(Some(HashAssertNode {
115 n_bytes,
116 preimage_digest,
117 expected_digest,
118 }))
119 }
120
121 fn tag(args: [Felt; 3]) -> Tag {
122 Tag::precompile(Self::id(), args).expect("hash precompile id is not framework-reserved")
123 }
124
125 fn digest_chunks() -> usize {
126 H::DIGEST_FELTS.div_ceil(8)
127 }
128}
129
130impl<H: HashFunction> Precompile for HashPrecompile<H> {
131 fn name(&self) -> &'static str {
132 H::NAME
133 }
134
135 fn id(&self) -> Felt {
136 Self::id()
137 }
138
139 fn decode(&self, args: [Felt; 3]) -> Option<NodeType> {
140 let disc = u32::try_from(args[0].as_canonical_u64()).ok()?;
141 if disc != ASSERT_DISC || args[2] != ZERO {
142 return None;
143 }
144 u32::try_from(args[1].as_canonical_u64()).ok()?;
145 Some(NodeType::Join)
146 }
147
148 fn evaluate(
149 &self,
150 args: [Felt; 3],
151 payload: &Payload,
152 context: &mut DeferredContext<'_>,
153 ) -> Result<Node, PrecompileError> {
154 let disc =
155 u32::try_from(args[0].as_canonical_u64()).map_err(|_| PrecompileError::InvalidNode)?;
156 let n_bytes =
157 u32::try_from(args[1].as_canonical_u64()).map_err(|_| PrecompileError::InvalidNode)?;
158 if disc != ASSERT_DISC || args[2] != ZERO {
159 return Err(PrecompileError::InvalidNode);
160 }
161
162 let (preimage_digest, expected_digest) = payload.as_join()?;
163 let preimage = chunks_child_to_bytes(
164 context,
165 preimage_digest,
166 n_chunks(n_bytes).get() as usize,
167 n_bytes as usize,
168 )?;
169 let expected = chunks_child_to_bytes(
170 context,
171 expected_digest,
172 Self::digest_chunks(),
173 H::DIGEST_FELTS * size_of::<u32>(),
174 )?;
175
176 if expected != H::hash(&preimage) {
177 return Err(PrecompileError::AssertionFailed);
178 }
179 Ok(Node::TRUE)
180 }
181}
182
183fn chunks_child_to_bytes(
184 context: &mut DeferredContext<'_>,
185 digest: Digest,
186 expected_chunks: usize,
187 n_bytes: usize,
188) -> Result<Vec<u8>, PrecompileError> {
189 let canonical_digest = context.evaluate_digest(digest)?;
190 let canonical_node = context.get_node(&canonical_digest).ok_or(PrecompileError::InvalidNode)?;
191 if canonical_node.tag() != Tag::CHUNKS {
192 return Err(PrecompileError::InvalidNode);
193 }
194 let chunks = canonical_node.payload().as_data()?;
195 chunks_to_bytes_exact(chunks, expected_chunks, n_bytes)
196}
197
198#[cfg(test)]
203pub(crate) fn assert_hash_precompile<H: HashFunction>() {
204 use alloc::{sync::Arc, vec, vec::Vec};
205
206 use miden_core::{
207 deferred::{DeferredState, PrecompileRegistry, TRUE_DIGEST, Tag, WireEntry},
208 utils::bytes_to_packed_u32_elements,
209 };
210
211 fn chunks_from_bytes(bytes: &[u8]) -> Vec<[Felt; 8]> {
212 Node::chunks_from_bytes(bytes)
213 .payload()
214 .as_data()
215 .expect("chunks_from_bytes creates data payload")
216 .to_vec()
217 }
218
219 fn digest_chunks<H: HashFunction>(input: &[u8]) -> Vec<[Felt; 8]> {
220 let mut felts = bytes_to_packed_u32_elements(&H::hash(input));
221 felts.resize(HashPrecompile::<H>::digest_chunks() * 8, ZERO);
222 felts
223 .as_chunks::<8>()
224 .0
225 .iter()
226 .map(|c| core::array::from_fn(|i| c[i]))
227 .collect()
228 }
229
230 let fresh = || {
231 DeferredState::new(
232 Arc::new(PrecompileRegistry::new().with_precompile(HashPrecompile::<H>::default())),
233 usize::MAX,
234 )
235 .expect("hash precompile initialization should fit the test budget")
236 };
237 let assert_registers = |state: &mut DeferredState,
238 n_bytes: u32,
239 preimage_chunks: Vec<[Felt; 8]>,
240 expected_chunks: Vec<[Felt; 8]>|
241 -> Result<Digest, PrecompileError> {
242 let preimage = state.register(Node::chunks(preimage_chunks).expect("preimage chunks"))?;
243 let expected = state.register(Node::chunks(expected_chunks).expect("expected chunks"))?;
244 state.register(HashPrecompile::<H>::assert_node(n_bytes, preimage, expected))
245 };
246 let assert_error = |err: PrecompileError, expected: PrecompileError| {
247 assert!(
248 matches!(
249 (err.root(), &expected),
250 (PrecompileError::InvalidNode, PrecompileError::InvalidNode)
251 | (PrecompileError::AssertionFailed, PrecompileError::AssertionFailed)
252 ),
253 "unexpected error root: {err:?}"
254 );
255 };
256
257 let pc = HashPrecompile::<H>::default();
258 assert_eq!(
259 pc.decode([Felt::from_u32(HashPrecompile::<H>::ASSERT_TAG_ID), Felt::from_u32(65), ZERO]),
260 Some(NodeType::Join),
261 );
262 assert!(pc.decode([Felt::from_u32(1), ZERO, ZERO]).is_none());
263 assert!(
264 pc.decode([
265 Felt::from_u32(HashPrecompile::<H>::ASSERT_TAG_ID),
266 Felt::from_u32(65),
267 Felt::from_u32(1),
268 ])
269 .is_none()
270 );
271 let non_u32 = Felt::new_unchecked(u64::from(u32::MAX) + 1);
272 assert!(pc.decode([non_u32, ZERO, ZERO]).is_none());
273 assert!(
274 pc.decode([Felt::from_u32(HashPrecompile::<H>::ASSERT_TAG_ID), non_u32, ZERO])
275 .is_none()
276 );
277
278 let assert_tag = HashPrecompile::<H>::assert_tag(65);
279 assert_eq!(HashPrecompile::<H>::decode_assert_tag(assert_tag).unwrap(), Some(65));
280 assert_eq!(HashPrecompile::<H>::decode_assert_tag(Tag::CHUNKS).unwrap(), None);
281 let invalid_assert_tag =
282 HashPrecompile::<H>::tag([Felt::from_u32(1), Felt::from_u32(65), ZERO]);
283 assert!(matches!(
284 HashPrecompile::<H>::decode_assert_tag(invalid_assert_tag),
285 Err(PrecompileError::InvalidNode)
286 ));
287
288 let assert_node = HashPrecompile::<H>::assert_node(65, TRUE_DIGEST, TRUE_DIGEST);
289 assert_eq!(HashPrecompile::<H>::decode_assert_node(&Node::TRUE).unwrap(), None);
290 assert_eq!(
291 HashPrecompile::<H>::decode_assert_node(&assert_node).unwrap(),
292 Some(HashAssertNode {
293 n_bytes: 65,
294 preimage_digest: TRUE_DIGEST,
295 expected_digest: TRUE_DIGEST,
296 })
297 );
298 let invalid_node = Node::join(invalid_assert_tag, TRUE_DIGEST, TRUE_DIGEST).unwrap();
299 assert!(matches!(
300 HashPrecompile::<H>::decode_assert_node(&invalid_node),
301 Err(PrecompileError::InvalidNode)
302 ));
303 let invalid_shape = Node::value(assert_tag, [ZERO; 8]).unwrap();
304 assert!(HashPrecompile::<H>::decode_assert_node(&invalid_shape).is_err());
305
306 let input = b"hash assertions consume generic chunks";
307 let mut state = fresh();
308 let assertion = assert_registers(
309 &mut state,
310 input.len() as u32,
311 chunks_from_bytes(input),
312 digest_chunks::<H>(input),
313 )
314 .expect("matching hash assertion should register");
315 assert_eq!(state.evaluate_digest(assertion).unwrap(), TRUE_DIGEST);
316 state.log_statement(assertion).expect("true assertion should log");
317
318 let mut wrong = digest_chunks::<H>(input);
319 wrong[0][0] = if wrong[0][0] == ZERO { Felt::from_u32(1) } else { ZERO };
320 let mut state = fresh();
321 let err = assert_registers(&mut state, input.len() as u32, chunks_from_bytes(input), wrong)
322 .unwrap_err();
323 assert_error(err, PrecompileError::AssertionFailed);
324
325 let too_long: Vec<u8> = (0u8..33).collect();
326 let mut state = fresh();
327 let err = assert_registers(
328 &mut state,
329 too_long.len() as u32,
330 vec![chunks_from_bytes(&too_long)[0]],
331 digest_chunks::<H>(&too_long),
332 )
333 .unwrap_err();
334 assert_error(err, PrecompileError::InvalidNode);
335
336 let mut padded = chunks_from_bytes(&[1, 2, 3]);
337 padded[0][0] = Felt::from_u32(u32::from_le_bytes([1, 2, 3, 0xaa]));
338 let mut state = fresh();
339 let err = assert_registers(&mut state, 3, padded, digest_chunks::<H>(&[1, 2, 3])).unwrap_err();
340 assert_error(err, PrecompileError::InvalidNode);
341
342 let non_u32 = Felt::new_unchecked(u64::from(u32::MAX) + 1);
343 let mut preimage = chunks_from_bytes(input);
344 preimage[0][0] = non_u32;
345 let mut state = fresh();
346 let err = assert_registers(&mut state, input.len() as u32, preimage, digest_chunks::<H>(input))
347 .unwrap_err();
348 assert_error(err, PrecompileError::InvalidNode);
349
350 let mut expected = digest_chunks::<H>(input);
351 expected[0][0] = non_u32;
352 let mut state = fresh();
353 let err = assert_registers(&mut state, input.len() as u32, chunks_from_bytes(input), expected)
354 .unwrap_err();
355 assert_error(err, PrecompileError::InvalidNode);
356
357 let precompile_owned_data = Node::try_data(
358 HashPrecompile::<H>::assert_tag(input.len() as u32),
359 chunks_from_bytes(input),
360 )
361 .expect("data node is syntactically constructible");
362 let mut state = fresh();
363 let preimage = state.register(precompile_owned_data).unwrap_err();
364 assert_error(preimage, PrecompileError::InvalidNode);
365
366 let mut state = fresh();
367 let zero = assert_registers(&mut state, 0, vec![[ZERO; 8]], digest_chunks::<H>(&[]))
368 .expect("zero-byte hash assertion should register");
369 assert_eq!(state.evaluate_digest(zero).unwrap(), TRUE_DIGEST);
370
371 let mut state = fresh();
372 let preimage_chunks = chunks_from_bytes(input);
373 let expected_chunks = digest_chunks::<H>(input);
374 let preimage = state.register(Node::chunks(preimage_chunks).unwrap()).unwrap();
375 let expected = state.register(Node::chunks(expected_chunks).unwrap()).unwrap();
376 let assertion_node = HashPrecompile::<H>::assert_node(input.len() as u32, preimage, expected);
377 let assertion = state.register(assertion_node).unwrap();
378 state.log_statement(assertion).unwrap();
379 let wire = state.to_wire().expect("hash assertion state should encode");
380 assert!(wire.entries.iter().any(|entry| matches!(
381 entry,
382 WireEntry::Join { tag, .. } if *tag == HashPrecompile::<H>::assert_tag(input.len() as u32)
383 )));
384 let mut rehydrated = DeferredState::from_wire(
385 Arc::new(PrecompileRegistry::new().with_precompile(HashPrecompile::<H>::default())),
386 &wire,
387 usize::MAX,
388 )
389 .expect("wire should rehydrate under the hash registry");
390 assert_eq!(rehydrated.evaluate_digest(rehydrated.root()).unwrap(), TRUE_DIGEST);
391}