openvm_keccak256_circuit/xorin/
trace.rs

1use std::{
2    borrow::BorrowMut,
3    mem::{align_of, size_of},
4};
5
6use openvm_circuit::{
7    arch::*,
8    system::memory::{
9        offline_checker::{MemoryReadAuxRecord, MemoryWriteBytesAuxRecord},
10        online::TracingMemory,
11        MemoryAuxColsFactory,
12    },
13};
14use openvm_circuit_primitives::AlignedBytesBorrow;
15use openvm_instructions::{
16    instruction::Instruction,
17    program::DEFAULT_PC_STEP,
18    riscv::{RV32_CELL_BITS, RV32_MEMORY_AS, RV32_REGISTER_AS, RV32_REGISTER_NUM_LIMBS},
19};
20use openvm_keccak256_transpiler::XorinOpcode;
21use openvm_rv32im_circuit::adapters::{read_rv32_register, tracing_read, tracing_write};
22use openvm_stark_backend::p3_field::PrimeField32;
23
24use crate::xorin::{columns::XorinVmCols, XorinVmExecutor, XorinVmFiller};
25
26#[derive(Clone, Copy)]
27pub struct XorinVmMetadata {}
28
29impl MultiRowMetadata for XorinVmMetadata {
30    fn get_num_rows(&self) -> usize {
31        1
32    }
33}
34
35pub(crate) type XorinVmRecordLayout = MultiRowLayout<XorinVmMetadata>;
36
37#[repr(C)]
38#[derive(AlignedBytesBorrow, Debug, Clone)]
39pub struct XorinVmRecordHeader {
40    pub from_pc: u32,
41    pub timestamp: u32,
42    pub rd_ptr: u32,
43    pub rs1_ptr: u32,
44    pub rs2_ptr: u32,
45    pub buffer: u32,
46    pub input: u32,
47    pub len: u32,
48    pub buffer_limbs: [u8; 136],
49    pub input_limbs: [u8; 136],
50    pub register_aux_cols: [MemoryReadAuxRecord; 3],
51    pub input_read_aux_cols: [MemoryReadAuxRecord; 34],
52    pub buffer_read_aux_cols: [MemoryReadAuxRecord; 34],
53    pub buffer_write_aux_cols: [MemoryWriteBytesAuxRecord<4>; 34],
54}
55
56pub struct XorinVmRecordMut<'a> {
57    pub inner: &'a mut XorinVmRecordHeader,
58}
59
60// Custom borrowing to split the buffer into a fixed `XorinVmRecord` header
61impl<'a> CustomBorrow<'a, XorinVmRecordMut<'a>, XorinVmRecordLayout> for [u8] {
62    fn custom_borrow(&'a mut self, _layout: XorinVmRecordLayout) -> XorinVmRecordMut<'a> {
63        let (record_buf, _rest) =
64            unsafe { self.split_at_mut_unchecked(size_of::<XorinVmRecordHeader>()) };
65        XorinVmRecordMut {
66            inner: record_buf.borrow_mut(),
67        }
68    }
69
70    unsafe fn extract_layout(&self) -> XorinVmRecordLayout {
71        XorinVmRecordLayout {
72            metadata: XorinVmMetadata {},
73        }
74    }
75}
76
77impl SizedRecord<XorinVmRecordLayout> for XorinVmRecordMut<'_> {
78    fn size(_layout: &XorinVmRecordLayout) -> usize {
79        size_of::<XorinVmRecordHeader>()
80    }
81
82    fn alignment(_layout: &XorinVmRecordLayout) -> usize {
83        align_of::<XorinVmRecordHeader>()
84    }
85}
86
87impl<F, RA> PreflightExecutor<F, RA> for XorinVmExecutor
88where
89    F: PrimeField32,
90    for<'buf> RA: RecordArena<'buf, XorinVmRecordLayout, XorinVmRecordMut<'buf>>,
91{
92    fn get_opcode_name(&self, _: usize) -> String {
93        format!("{:?}", XorinOpcode::XORIN)
94    }
95
96    fn execute(
97        &self,
98        state: VmStateMut<F, TracingMemory, RA>,
99        instruction: &Instruction<F>,
100    ) -> Result<(), ExecutionError> {
101        let &Instruction { a, b, c, .. } = instruction;
102
103        // Reading the length first without tracing to allocate a record of correct size
104        let guest_mem = state.memory.data();
105        let len = read_rv32_register(guest_mem, c.as_canonical_u32()) as usize;
106        // Safety: length has to be multiple of 4
107        // This is enforced by how the guest program calls the xorin opcode
108        // Xorin opcode is only called through the keccak update guest program
109        debug_assert!(len.is_multiple_of(4));
110        let num_reads = len.div_ceil(4);
111
112        // safety: the below alloc uses MultiRowLayout alloc implementation because
113        // XorinVmRecordLayout is a MultiRowLayout since get_num_rows() = 1, this will
114        // alloc_buffer of size width where width is the width of the trace matrix
115        // then it takes a prefix of this allocated buffer through custom borrow
116        // of length XorinVmRecordLayout size and return it as the below `record`
117        let record = state
118            .ctx
119            .alloc(XorinVmRecordLayout::new(XorinVmMetadata {}));
120
121        record.inner.from_pc = *state.pc;
122        record.inner.timestamp = state.memory.timestamp();
123        record.inner.rd_ptr = a.as_canonical_u32();
124        record.inner.rs1_ptr = b.as_canonical_u32();
125        record.inner.rs2_ptr = c.as_canonical_u32();
126
127        record.inner.buffer = u32::from_le_bytes(tracing_read(
128            state.memory,
129            RV32_REGISTER_AS,
130            record.inner.rd_ptr,
131            &mut record.inner.register_aux_cols[0].prev_timestamp,
132        ));
133
134        record.inner.input = u32::from_le_bytes(tracing_read(
135            state.memory,
136            RV32_REGISTER_AS,
137            record.inner.rs1_ptr,
138            &mut record.inner.register_aux_cols[1].prev_timestamp,
139        ));
140
141        record.inner.len = u32::from_le_bytes(tracing_read(
142            state.memory,
143            RV32_REGISTER_AS,
144            record.inner.rs2_ptr,
145            &mut record.inner.register_aux_cols[2].prev_timestamp,
146        ));
147
148        debug_assert!(record.inner.buffer as usize + len <= (1 << self.pointer_max_bits));
149        debug_assert!(record.inner.input as usize + len < (1 << self.pointer_max_bits));
150        debug_assert!(record.inner.len < (1 << self.pointer_max_bits));
151
152        // read buffer
153        for idx in 0..num_reads {
154            let read = tracing_read::<4>(
155                state.memory,
156                RV32_MEMORY_AS,
157                record.inner.buffer + (idx * 4) as u32,
158                &mut record.inner.buffer_read_aux_cols[idx].prev_timestamp,
159            );
160            record.inner.buffer_limbs[4 * idx..4 * (idx + 1)].copy_from_slice(&read);
161        }
162
163        // read input
164        for idx in 0..num_reads {
165            let read = tracing_read::<4>(
166                state.memory,
167                RV32_MEMORY_AS,
168                record.inner.input + (idx * 4) as u32,
169                &mut record.inner.input_read_aux_cols[idx].prev_timestamp,
170            );
171            record.inner.input_limbs[4 * idx..4 * (idx + 1)].copy_from_slice(&read);
172        }
173
174        let mut result = [0u8; 136];
175
176        // execute xorin
177        for ((x_xor_y, &x), &y) in result
178            .iter_mut()
179            .zip(record.inner.buffer_limbs.iter())
180            .zip(record.inner.input_limbs.iter())
181        {
182            *x_xor_y = x ^ y;
183        }
184
185        // write result
186        for idx in 0..num_reads {
187            let mut word: [u8; 4] = [0u8; 4];
188            word.copy_from_slice(&result[4 * idx..4 * (idx + 1)]);
189            tracing_write(
190                state.memory,
191                RV32_MEMORY_AS,
192                record.inner.buffer + (idx * 4) as u32,
193                word,
194                &mut record.inner.buffer_write_aux_cols[idx].prev_timestamp,
195                &mut record.inner.buffer_write_aux_cols[idx].prev_data,
196            );
197        }
198
199        *state.pc = state.pc.wrapping_add(DEFAULT_PC_STEP);
200
201        Ok(())
202    }
203}
204
205impl<F: PrimeField32> TraceFiller<F> for XorinVmFiller {
206    fn fill_trace_row(&self, mem_helper: &MemoryAuxColsFactory<F>, mut row_slice: &mut [F]) {
207        let record: XorinVmRecordMut = unsafe {
208            get_record_from_slice(
209                &mut row_slice,
210                XorinVmRecordLayout {
211                    metadata: XorinVmMetadata {},
212                },
213            )
214        };
215
216        // Safety: the clone here is necessary because the XorinVmCols uses the same buffer
217        let record = record.inner.clone();
218        row_slice.fill(F::ZERO);
219        let trace_row: &mut XorinVmCols<F> = row_slice.borrow_mut();
220
221        trace_row.instruction.pc = F::from_u32(record.from_pc);
222        trace_row.instruction.is_enabled = F::ONE;
223        trace_row.instruction.buffer_reg_ptr = F::from_u32(record.rd_ptr);
224        trace_row.instruction.input_reg_ptr = F::from_u32(record.rs1_ptr);
225        trace_row.instruction.len_reg_ptr = F::from_u32(record.rs2_ptr);
226        trace_row.instruction.buffer_ptr = F::from_u32(record.buffer);
227        let buffer_ptr_u8: [u8; 4] = record.buffer.to_le_bytes();
228        let buffer_ptr_limbs: [F; 4] = [
229            F::from_u8(buffer_ptr_u8[0]),
230            F::from_u8(buffer_ptr_u8[1]),
231            F::from_u8(buffer_ptr_u8[2]),
232            F::from_u8(buffer_ptr_u8[3]),
233        ];
234        trace_row.instruction.buffer_ptr_limbs = buffer_ptr_limbs;
235        trace_row.instruction.input_ptr = F::from_u32(record.input);
236        let input_ptr_u8: [u8; 4] = record.input.to_le_bytes();
237        let input_ptr_limbs: [F; 4] = [
238            F::from_u8(input_ptr_u8[0]),
239            F::from_u8(input_ptr_u8[1]),
240            F::from_u8(input_ptr_u8[2]),
241            F::from_u8(input_ptr_u8[3]),
242        ];
243        trace_row.instruction.input_ptr_limbs = input_ptr_limbs;
244        trace_row.instruction.len = F::from_u32(record.len);
245        let len_u8: [u8; 4] = record.len.to_le_bytes();
246        let len_limbs: [F; 4] = [
247            F::from_u8(len_u8[0]),
248            F::from_u8(len_u8[1]),
249            F::from_u8(len_u8[2]),
250            F::from_u8(len_u8[3]),
251        ];
252        trace_row.instruction.len_limbs = len_limbs;
253        trace_row.instruction.start_timestamp = F::from_u32(record.timestamp);
254
255        for i in 0..(record.len / 4) {
256            trace_row.sponge.is_padding_bytes[i as usize] = F::ZERO;
257        }
258        for i in (record.len / 4)..34 {
259            trace_row.sponge.is_padding_bytes[i as usize] = F::ONE;
260        }
261
262        let mut timestamp = record.timestamp;
263        let record_len: usize = record.len as usize;
264        let num_reads: usize = record_len.div_ceil(4);
265
266        for t in 0..3 {
267            mem_helper.fill(
268                record.register_aux_cols[t].prev_timestamp,
269                timestamp,
270                trace_row.mem_oc.register_aux_cols[t].as_mut(),
271            );
272
273            timestamp += 1;
274        }
275
276        for t in 0..num_reads {
277            mem_helper.fill(
278                record.buffer_read_aux_cols[t].prev_timestamp,
279                timestamp,
280                trace_row.mem_oc.buffer_bytes_read_aux_cols[t].as_mut(),
281            );
282            timestamp += 1;
283        }
284
285        for t in 0..num_reads {
286            mem_helper.fill(
287                record.input_read_aux_cols[t].prev_timestamp,
288                timestamp,
289                trace_row.mem_oc.input_bytes_read_aux_cols[t].as_mut(),
290            );
291            timestamp += 1;
292        }
293
294        // safety note: we leave the upper record_len..134 bytes with zeroes
295        // because they are just padding bytes and unused by the chip
296        for i in 0..record_len {
297            trace_row.sponge.preimage_buffer_bytes[i] = F::from_u8(record.buffer_limbs[i]);
298            trace_row.sponge.input_bytes[i] = F::from_u8(record.input_limbs[i]);
299            trace_row.sponge.postimage_buffer_bytes[i] =
300                F::from_u8(record.buffer_limbs[i] ^ record.input_limbs[i]);
301            let b_val = record.buffer_limbs[i] as u32;
302            let c_val = record.input_limbs[i] as u32;
303            self.bitwise_lookup_chip.request_xor(b_val, c_val);
304        }
305
306        for t in 0..num_reads {
307            mem_helper.fill(
308                record.buffer_write_aux_cols[t].prev_timestamp,
309                timestamp,
310                trace_row.mem_oc.buffer_bytes_write_aux_cols[t].as_mut(),
311            );
312            trace_row.mem_oc.buffer_bytes_write_aux_cols[t].prev_data =
313                record.buffer_write_aux_cols[t].prev_data.map(F::from_u8);
314            timestamp += 1;
315        }
316
317        let buffer_ptr_limbs = record.buffer.to_le_bytes();
318        let input_ptr_limbs = record.input.to_le_bytes();
319        let len_limbs = record.len.to_le_bytes();
320
321        let need_range_check = [
322            buffer_ptr_limbs.last().unwrap(),
323            input_ptr_limbs.last().unwrap(),
324            len_limbs.last().unwrap(),
325            len_limbs.last().unwrap(),
326        ];
327
328        let limb_shift = 1 << (RV32_CELL_BITS * RV32_REGISTER_NUM_LIMBS - self.pointer_max_bits);
329
330        for pair in need_range_check.chunks_exact(2) {
331            self.bitwise_lookup_chip
332                .request_range((pair[0] * limb_shift) as u32, (pair[1] * limb_shift) as u32);
333        }
334    }
335}