openvm_keccak256_circuit/keccakf_op/
trace.rs1use core::convert::TryInto;
2use std::{
3 borrow::BorrowMut,
4 mem::{align_of, size_of},
5 sync::{Arc, Mutex},
6};
7
8use openvm_circuit::{
9 arch::*,
10 system::memory::{
11 offline_checker::MemoryReadAuxRecord, online::TracingMemory, MemoryAuxColsFactory,
12 SharedMemoryHelper,
13 },
14};
15use openvm_circuit_primitives::{
16 bitwise_op_lookup::SharedBitwiseOperationLookupChip, AlignedBytesBorrow, Chip,
17};
18use openvm_cpu_backend::CpuBackend;
19use openvm_instructions::{
20 instruction::Instruction,
21 program::DEFAULT_PC_STEP,
22 riscv::{RV32_CELL_BITS, RV32_MEMORY_AS, RV32_REGISTER_AS, RV32_REGISTER_NUM_LIMBS},
23};
24use openvm_keccak256_transpiler::KeccakfOpcode;
25use openvm_rv32im_circuit::adapters::{timed_write, tracing_read};
26use openvm_stark_backend::{
27 p3_field::PrimeField32,
28 p3_matrix::{dense::RowMajorMatrix, Matrix},
29 p3_maybe_rayon::prelude::*,
30 prover::AirProvingContext,
31 StarkProtocolConfig, Val,
32};
33
34use super::{KeccakfExecutor, NUM_OP_ROWS_PER_INS};
35use crate::{
36 keccakf_op::{columns::KeccakfOpCols, keccakf_postimage_bytes},
37 KECCAK_WIDTH_BYTES, KECCAK_WIDTH_WORDS, KECCAK_WORD_SIZE,
38};
39
40#[derive(derive_new::new)]
41pub struct KeccakfOpChip<F> {
42 pub bitwise_lookup_chip: SharedBitwiseOperationLookupChip<8>,
43 pub pointer_max_bits: usize,
44 pub mem_helper: SharedMemoryHelper<F>,
45 pub shared_records: Arc<Mutex<Vec<KeccakfRecord>>>,
49}
50
51impl<SC, RA> Chip<RA, CpuBackend<SC>> for KeccakfOpChip<Val<SC>>
52where
53 SC: StarkProtocolConfig,
54 Val<SC>: PrimeField32,
55 RA: RowMajorMatrixArena<Val<SC>>,
56{
57 fn generate_proving_ctx(&self, arena: RA) -> AirProvingContext<CpuBackend<SC>> {
58 let rows_used = arena.trace_offset() / arena.width();
59 let mut trace = arena.into_matrix();
60 let mem_helper = self.mem_helper.as_borrowed();
61 self.fill_trace(&mem_helper, &mut trace, rows_used);
62 AirProvingContext::simple_no_pis(trace)
63 }
64}
65
66#[derive(Clone, Copy, Default)]
67pub struct KeccakfMetadata;
68
69impl MultiRowMetadata for KeccakfMetadata {
70 fn get_num_rows(&self) -> usize {
71 NUM_OP_ROWS_PER_INS
72 }
73}
74
75pub(crate) type KeccakfRecordLayout = MultiRowLayout<KeccakfMetadata>;
76
77#[repr(C)]
78#[derive(AlignedBytesBorrow, Debug, Clone)]
79pub struct KeccakfRecord {
80 pub pc: u32,
81 pub timestamp: u32,
82 pub rd_ptr: u32,
83 pub buffer_ptr: u32,
84 pub rd_aux: MemoryReadAuxRecord,
85 pub buffer_word_aux: [MemoryReadAuxRecord; KECCAK_WIDTH_WORDS],
86 pub preimage_buffer_bytes: [u8; KECCAK_WIDTH_BYTES],
87}
88
89pub struct KeccakfRecordMut<'a> {
91 pub inner: &'a mut KeccakfRecord,
92}
93
94impl<'a> CustomBorrow<'a, KeccakfRecordMut<'a>, KeccakfRecordLayout> for [u8] {
95 fn custom_borrow(&'a mut self, _layout: KeccakfRecordLayout) -> KeccakfRecordMut<'a> {
96 let (record_buf, _rest) =
97 unsafe { self.split_at_mut_unchecked(size_of::<KeccakfRecord>()) };
98 KeccakfRecordMut {
99 inner: record_buf.borrow_mut(),
100 }
101 }
102
103 unsafe fn extract_layout(&self) -> KeccakfRecordLayout {
104 KeccakfRecordLayout::new(KeccakfMetadata)
105 }
106}
107
108impl SizedRecord<KeccakfRecordLayout> for KeccakfRecordMut<'_> {
109 fn size(_layout: &KeccakfRecordLayout) -> usize {
110 size_of::<KeccakfRecord>()
111 }
112
113 fn alignment(_layout: &KeccakfRecordLayout) -> usize {
114 align_of::<KeccakfRecord>()
115 }
116}
117
118impl<F, RA> PreflightExecutor<F, RA> for KeccakfExecutor
119where
120 F: PrimeField32,
121 for<'buf> RA: RecordArena<'buf, KeccakfRecordLayout, &'buf mut KeccakfRecord>,
122{
123 fn get_opcode_name(&self, _: usize) -> String {
124 format!("{:?}", KeccakfOpcode::KECCAKF)
125 }
126
127 fn execute(
128 &self,
129 state: VmStateMut<F, TracingMemory, RA>,
130 instruction: &Instruction<F>,
131 ) -> Result<(), ExecutionError> {
132 let &Instruction { a, .. } = instruction;
133 let rd_ptr = a.as_canonical_u32();
134
135 let record = state.ctx.alloc(KeccakfRecordLayout::new(KeccakfMetadata));
136
137 record.pc = *state.pc;
138 record.timestamp = state.memory.timestamp();
139 record.rd_ptr = rd_ptr;
140 let buffer_ptr = u32::from_le_bytes(tracing_read(
141 state.memory,
142 RV32_REGISTER_AS,
143 rd_ptr,
144 &mut record.rd_aux.prev_timestamp,
145 ));
146 record.buffer_ptr = buffer_ptr;
147
148 let guest_mem = state.memory.data();
149 let prestate =
153 unsafe { guest_mem.get_slice(RV32_MEMORY_AS, record.buffer_ptr, KECCAK_WIDTH_BYTES) };
154 record.preimage_buffer_bytes.copy_from_slice(prestate);
155 let poststate = keccakf_postimage_bytes(&record.preimage_buffer_bytes);
156 for (word_idx, (word, aux)) in poststate
157 .chunks_exact(KECCAK_WORD_SIZE)
158 .zip(&mut record.buffer_word_aux)
159 .enumerate()
160 {
161 let (t_prev, _) = timed_write::<KECCAK_WORD_SIZE>(
163 state.memory,
164 RV32_MEMORY_AS,
165 buffer_ptr + (word_idx * KECCAK_WORD_SIZE) as u32,
166 word.try_into().unwrap(),
167 );
168 aux.prev_timestamp = t_prev;
169 }
170
171 *state.pc = state.pc.wrapping_add(DEFAULT_PC_STEP);
172 Ok(())
173 }
174}
175
176impl<F: PrimeField32> TraceFiller<F> for KeccakfOpChip<F> {
177 fn fill_trace(
178 &self,
179 mem_helper: &MemoryAuxColsFactory<F>,
180 trace_matrix: &mut RowMajorMatrix<F>,
181 rows_used: usize,
182 ) {
183 if rows_used == 0 {
184 return;
185 }
186 assert!(rows_used.is_multiple_of(NUM_OP_ROWS_PER_INS));
187
188 let width = trace_matrix.width();
189 let (trace, dummy_trace) = trace_matrix.values.split_at_mut(rows_used * width);
190 let records = trace
193 .par_chunks_exact_mut(width * NUM_OP_ROWS_PER_INS)
194 .map(|mut row| {
195 let record: &mut KeccakfRecord = unsafe {
196 get_record_from_slice(&mut row, KeccakfRecordLayout::new(KeccakfMetadata))
197 };
198 record.clone()
199 })
200 .collect::<Vec<_>>();
201 dummy_trace.fill(F::ZERO);
202
203 trace
204 .par_chunks_exact_mut(width * NUM_OP_ROWS_PER_INS)
205 .zip(records.par_iter())
206 .for_each(|(row, record)| {
207 row.fill(F::ZERO);
208
209 let postimage_buffer_bytes = keccakf_postimage_bytes(&record.preimage_buffer_bytes);
210 let buffer_ptr_limbs = record.buffer_ptr.to_le_bytes();
211
212 let local: &mut KeccakfOpCols<F> = row.borrow_mut();
213
214 local.pc = F::from_u32(record.pc);
215 local.is_valid = F::ONE;
216 local.timestamp = F::from_u32(record.timestamp);
217 local.rd_ptr = F::from_u32(record.rd_ptr);
218 local.buffer_ptr_limbs = buffer_ptr_limbs.map(F::from_u8);
219
220 for (dst, &byte) in local.preimage.iter_mut().zip(&record.preimage_buffer_bytes) {
221 *dst = F::from_u8(byte);
222 }
223 for (dst, &byte) in local.postimage.iter_mut().zip(&postimage_buffer_bytes) {
224 *dst = F::from_u8(byte);
225 }
226
227 let mut timestamp = record.timestamp;
228 mem_helper.fill(
229 record.rd_aux.prev_timestamp,
230 record.timestamp,
231 local.rd_aux.as_mut(),
232 );
233 timestamp += 1;
234 for (aux, record_aux) in local
235 .buffer_word_aux
236 .iter_mut()
237 .zip(&record.buffer_word_aux)
238 {
239 mem_helper.fill(record_aux.prev_timestamp, timestamp, aux);
240 timestamp += 1;
241 }
242
243 let limb_shift = 1u32
244 << (RV32_CELL_BITS * RV32_REGISTER_NUM_LIMBS - self.pointer_max_bits) as u32;
245 let scaled_limb =
246 (buffer_ptr_limbs[RV32_REGISTER_NUM_LIMBS - 1] as u32) * limb_shift;
247 self.bitwise_lookup_chip
248 .request_range(scaled_limb, scaled_limb);
249
250 for pair in postimage_buffer_bytes.chunks_exact(2) {
251 self.bitwise_lookup_chip
252 .request_range(pair[0] as u32, pair[1] as u32);
253 }
254 });
255 *self.shared_records.lock().unwrap() = records;
256 }
257}