openvm_keccak256_circuit/xorin/
trace.rs1use 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
60impl<'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 let guest_mem = state.memory.data();
105 let len = read_rv32_register(guest_mem, c.as_canonical_u32()) as usize;
106 debug_assert!(len.is_multiple_of(4));
110 let num_reads = len.div_ceil(4);
111
112 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 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 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 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 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 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 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}