@@ -5,6 +5,114 @@ use lean_vm::*;
55use rand:: { RngExt , SeedableRng , rngs:: StdRng } ;
66use utils:: { init_tracing, poseidon16_compress} ;
77
8+ #[ test]
9+ #[ ignore = "benchmark; run with `cargo test --release -p lean_prover bench_poseidon -- --ignored --nocapture`" ]
10+ fn bench_poseidon ( ) {
11+ utils:: init_tracing ( ) ;
12+ let n_poseidon_calls = std:: env:: var ( "POSEIDON_BENCH_CALLS" )
13+ . ok ( )
14+ . map ( |raw| raw. parse :: < usize > ( ) . expect ( "POSEIDON_BENCH_CALLS must be a usize" ) )
15+ . unwrap_or ( 1 ) ;
16+ let program_str = format ! (
17+ r#"
18+ N_POSEIDON_CALLS = {n_poseidon_calls}
19+ DIGEST_LEN = 8
20+
21+ def main():
22+ input_left = 0
23+ input_right = DIGEST_LEN
24+ outputs = Array(N_POSEIDON_CALLS * DIGEST_LEN)
25+ for i in dynamic_unroll(0, N_POSEIDON_CALLS, 20):
26+ out = outputs + i * DIGEST_LEN
27+ poseidon16_compress(input_left, input_right, out)
28+ return
29+ "#
30+ ) ;
31+
32+ let public_input: Vec < F > = ( 0 ..16 ) . map ( F :: new) . collect ( ) ;
33+ let bytecode = compile_program ( & ProgramSource :: Raw ( program_str) ) ;
34+ let witness = ExecutionWitness :: default ( ) ;
35+ let starting_log_inv_rate = 1 ;
36+
37+ let time = std:: time:: Instant :: now ( ) ;
38+ let proof = prove_execution (
39+ & bytecode,
40+ & public_input,
41+ & witness,
42+ & default_whir_config ( starting_log_inv_rate) ,
43+ false ,
44+ ) ;
45+ let proof_time = time. elapsed ( ) ;
46+ let proof_size_kib = proof. proof . proof_size_fe ( ) * F :: bits ( ) / ( 8 * 1024 ) ;
47+
48+ println ! ( "{}" , proof. metadata. display( ) ) ;
49+ println ! ( "Proof time: {:.3} s" , proof_time. as_secs_f32( ) ) ;
50+ println ! ( "Proof size: {proof_size_kib} KiB" ) ;
51+
52+ verify_execution ( & bytecode, & public_input, proof. proof ) . unwrap ( ) ;
53+ }
54+
55+ #[ test]
56+ #[ ignore = "benchmark; run with `cargo test --release -p lean_prover bench_sha256_compress -- --ignored --nocapture`" ]
57+ fn bench_sha256_compress ( ) {
58+ utils:: init_tracing ( ) ;
59+ let n_sha_calls = std:: env:: var ( "SHA256_BENCH_CALLS" )
60+ . ok ( )
61+ . map ( |raw| raw. parse :: < usize > ( ) . expect ( "SHA256_BENCH_CALLS must be a usize" ) )
62+ . unwrap_or ( 1 ) ;
63+ const SHA_FIXTURE_STRIDE : usize = SHA256_STATE_LIMBS + SHA256_BLOCK_LIMBS + SHA256_STATE_LIMBS ;
64+ let program_str = format ! (
65+ r#"
66+ N_SHA_CALLS = {n_sha_calls}
67+ SHA_FIXTURE_STRIDE = 64
68+
69+ def main():
70+ for j in unroll(0, N_SHA_CALLS):
71+ base = j * SHA_FIXTURE_STRIDE
72+ state = base
73+ block = base + 16
74+ expected = base + 48
75+ out = Array(16)
76+ sha256_compress(state, block, out)
77+
78+ for i in unroll(0, 16):
79+ assert out[i] == expected[i]
80+ return
81+ "#
82+ ) ;
83+
84+ let mut public_input = vec ! [ F :: ZERO ; n_sha_calls * SHA_FIXTURE_STRIDE ] ;
85+ let expected = words_to_field_limbs_le ( sha256_compress_words ( SHA256_IV , SHA256_ABC_BLOCK ) ) ;
86+ for j in 0 ..n_sha_calls {
87+ let base = j * SHA_FIXTURE_STRIDE ;
88+ public_input[ base..base + SHA256_STATE_LIMBS ] . copy_from_slice ( & words_to_field_limbs_le ( SHA256_IV ) ) ;
89+ public_input[ base + 16 ..base + 16 + SHA256_BLOCK_LIMBS ]
90+ . copy_from_slice ( & words_to_field_limbs_le ( SHA256_ABC_BLOCK ) ) ;
91+ public_input[ base + 48 ..base + 48 + SHA256_STATE_LIMBS ] . copy_from_slice ( & expected) ;
92+ }
93+
94+ let bytecode = compile_program ( & ProgramSource :: Raw ( program_str) ) ;
95+ let witness = ExecutionWitness :: default ( ) ;
96+ let starting_log_inv_rate = 1 ;
97+
98+ let time = std:: time:: Instant :: now ( ) ;
99+ let proof = prove_execution (
100+ & bytecode,
101+ & public_input,
102+ & witness,
103+ & default_whir_config ( starting_log_inv_rate) ,
104+ false ,
105+ ) ;
106+ let proof_time = time. elapsed ( ) ;
107+ let proof_size_kib = proof. proof . proof_size_fe ( ) * F :: bits ( ) / ( 8 * 1024 ) ;
108+
109+ println ! ( "{}" , proof. metadata. display( ) ) ;
110+ println ! ( "Proof time: {:.3} s" , proof_time. as_secs_f32( ) ) ;
111+ println ! ( "Proof size: {proof_size_kib} KiB" ) ;
112+
113+ verify_execution ( & bytecode, & public_input, proof. proof ) . unwrap ( ) ;
114+ }
115+
8116#[ test]
9117fn test_zk_vm_all_precompiles ( ) {
10118 let program_str = r#"
@@ -17,6 +125,15 @@ def main():
17125 pub_start = 0
18126 poseidon16_compress(pub_start + 4 * DIGEST_LEN, pub_start + 5 * DIGEST_LEN, pub_start + 6 * DIGEST_LEN)
19127
128+ # Keep the SHA fixture away from the extension-op fixture ranges below.
129+ sha_state = pub_start + 1400
130+ sha_block = sha_state + 16
131+ sha_expected = sha_block + 32
132+ sha_out = Array(16)
133+ sha256_compress(sha_state, sha_block, sha_out)
134+ for i in unroll(0, 16):
135+ assert sha_out[i] == sha_expected[i]
136+
20137 base_ptr = pub_start + 88
21138 ext_a_ptr = pub_start + 88 + N
22139 ext_b_ptr = pub_start + 88 + N * (DIM + 1)
@@ -62,6 +179,19 @@ def main():
62179 let poseidon_24_input: [ F ; 24 ] = rng. random ( ) ;
63180 public_input[ 56 ..80 ] . copy_from_slice ( & poseidon_24_input) ;
64181
182+ // SHA256 compression test data: IV + padded "abc" block.
183+ // This mirrors the program's pub_start + 1400 offset; public_input is 2^13 cells,
184+ // so the state, block, and expected digest all fit in the public memory prefix.
185+ let sha_state_ptr = 1400 ;
186+ let sha_block_ptr = sha_state_ptr + SHA256_STATE_LIMBS ;
187+ let sha_expected_ptr = sha_block_ptr + SHA256_BLOCK_LIMBS ;
188+ public_input[ sha_state_ptr..sha_state_ptr + SHA256_STATE_LIMBS ]
189+ . copy_from_slice ( & words_to_field_limbs_le ( SHA256_IV ) ) ;
190+ public_input[ sha_block_ptr..sha_block_ptr + SHA256_BLOCK_LIMBS ]
191+ . copy_from_slice ( & words_to_field_limbs_le ( SHA256_ABC_BLOCK ) ) ;
192+ let sha_expected = words_to_field_limbs_le ( sha256_compress_words ( SHA256_IV , SHA256_ABC_BLOCK ) ) ;
193+ public_input[ sha_expected_ptr..sha_expected_ptr + SHA256_STATE_LIMBS ] . copy_from_slice ( & sha_expected) ;
194+
65195 // Extension op operands: base[N], ext_a[N], ext_b[N]
66196 let base_slice: [ F ; N ] = rng. random ( ) ;
67197 let ext_a_slice: [ EF ; N ] = rng. random ( ) ;
0 commit comments