openvm_keccak256_circuit/keccakf_op/
trace.rs

1use 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    // NOTE[jpw]: this is an awkward way to pass data from this execution chip to the
46    // KeccakfPeriphery chip. This can be improved with a redesign of how record arenas are shared
47    // with chips.
48    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
89/// Mutable reference wrapper for KeccakfRecord, used for record seeking in CUDA tests
90pub 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        // SAFETY:
150        // - RV32_MEMORY_AS (2) consists of `u8`
151        // - get_slice will panic (if protected mode) if out of bounds
152        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            // We don't need prev_data since we read it earlier
162            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        // For clarity we just clone the records into a separate vector to avoid dealing with unsafe
191        // overwriting
192        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}