1use std::{
2 array::from_fn,
3 borrow::{Borrow, BorrowMut},
4 mem::{align_of, size_of},
5 sync::Arc,
6};
7
8use itertools::Itertools;
9use openvm_circuit::{
10 arch::{
11 get_record_from_slice, CustomBorrow, ExecutionError, MultiRowLayout, MultiRowMetadata,
12 PreflightExecutor, RecordArena, SizedRecord, TraceFiller, VmField, VmStateMut,
13 DEFAULT_BLOCK_SIZE,
14 },
15 system::memory::{
16 offline_checker::{MemoryReadAuxRecord, MemoryWriteBytesAuxRecord},
17 online::TracingMemory,
18 MemoryAuxColsFactory,
19 },
20};
21use openvm_circuit_primitives::{
22 bitwise_op_lookup::SharedBitwiseOperationLookupChip, AlignedBytesBorrow,
23};
24use openvm_deferral_transpiler::DeferralOpcode;
25use openvm_instructions::{
26 instruction::Instruction,
27 program::DEFAULT_PC_STEP,
28 riscv::{RV32_CELL_BITS, RV32_MEMORY_AS, RV32_REGISTER_AS, RV32_REGISTER_NUM_LIMBS},
29};
30use openvm_rv32im_circuit::adapters::{
31 memory_read, read_rv32_register, tracing_read, tracing_write,
32};
33use openvm_stark_backend::{p3_field::PrimeField32, p3_matrix::dense::RowMajorMatrix};
34use openvm_stark_sdk::config::baby_bear_poseidon2::DIGEST_SIZE;
35
36use crate::{
37 canonicity::CanonicityTraceGen,
38 count::DeferralCircuitCountChip,
39 output::DeferralOutputCols,
40 poseidon2::DeferralPoseidon2Chip,
41 utils::{
42 f_commit_to_bytes, join_memory_ops, memory_op_chunk, split_output, DIGEST_MEMORY_OPS,
43 F_NUM_BYTES, OUTPUT_TOTAL_BYTES, OUTPUT_TOTAL_MEMORY_OPS,
44 },
45};
46
47#[derive(Clone, Copy, Debug, Default)]
48pub struct DeferralOutputMetadata {
49 pub num_rows: usize,
50}
51
52impl MultiRowMetadata for DeferralOutputMetadata {
53 #[inline(always)]
54 fn get_num_rows(&self) -> usize {
55 self.num_rows
56 }
57}
58
59pub(crate) type DeferralOutputLayout = MultiRowLayout<DeferralOutputMetadata>;
60
61#[repr(C)]
62#[derive(AlignedBytesBorrow, Debug, Clone)]
63pub struct DeferralOutputRecordHeader {
64 pub from_pc: u32,
65 pub from_timestamp: u32,
66 pub rd_ptr: u32,
67 pub rs_ptr: u32,
68 pub deferral_idx: u32,
69 pub num_rows: u32,
70
71 pub rd_val: [u8; RV32_REGISTER_NUM_LIMBS],
73 pub rs_val: [u8; RV32_REGISTER_NUM_LIMBS],
74 pub rd_aux: MemoryReadAuxRecord,
75 pub rs_aux: MemoryReadAuxRecord,
76
77 pub output_commit_and_len_aux: [MemoryReadAuxRecord; OUTPUT_TOTAL_MEMORY_OPS],
79}
80
81pub struct DeferralOutputRecordMut<'a> {
82 pub header: &'a mut DeferralOutputRecordHeader,
83 pub write_bytes: &'a mut [u8],
84 pub write_aux: &'a mut [MemoryWriteBytesAuxRecord<DEFAULT_BLOCK_SIZE>],
85}
86
87impl<'a> CustomBorrow<'a, DeferralOutputRecordMut<'a>, DeferralOutputLayout> for [u8] {
88 fn custom_borrow(&'a mut self, layout: DeferralOutputLayout) -> DeferralOutputRecordMut<'a> {
89 let (header_buf, rest) =
92 unsafe { self.split_at_mut_unchecked(size_of::<DeferralOutputRecordHeader>()) };
93
94 let num_write_rows = layout.metadata.num_rows.saturating_sub(1);
98 let (write_bytes, rest) =
99 unsafe { rest.split_at_mut_unchecked(num_write_rows * DIGEST_SIZE) };
100
101 let (_, write_aux_buf, _) =
107 unsafe { rest.align_to_mut::<MemoryWriteBytesAuxRecord<DEFAULT_BLOCK_SIZE>>() };
108
109 DeferralOutputRecordMut {
110 header: header_buf.borrow_mut(),
111 write_bytes,
112 write_aux: &mut write_aux_buf[..num_write_rows * DIGEST_MEMORY_OPS],
113 }
114 }
115
116 unsafe fn extract_layout(&self) -> DeferralOutputLayout {
117 let record: &DeferralOutputRecordHeader = self.borrow();
118 DeferralOutputLayout {
119 metadata: DeferralOutputMetadata {
120 num_rows: record.num_rows as usize,
121 },
122 }
123 }
124}
125
126impl<'a> SizedRecord<DeferralOutputLayout> for DeferralOutputRecordMut<'a> {
127 fn size(layout: &DeferralOutputLayout) -> usize {
128 let mut total_len = size_of::<DeferralOutputRecordHeader>();
129 let num_write_rows = layout.metadata.num_rows.saturating_sub(1);
130 total_len += num_write_rows * DIGEST_SIZE;
131 total_len =
132 total_len.next_multiple_of(align_of::<MemoryWriteBytesAuxRecord<DEFAULT_BLOCK_SIZE>>());
133 total_len += num_write_rows
134 * DIGEST_MEMORY_OPS
135 * size_of::<MemoryWriteBytesAuxRecord<DEFAULT_BLOCK_SIZE>>();
136 total_len
137 }
138
139 fn alignment(_layout: &DeferralOutputLayout) -> usize {
140 align_of::<DeferralOutputRecordHeader>()
141 }
142}
143
144#[derive(Clone, Copy, Debug, derive_new::new)]
145pub struct DeferralOutputExecutor;
146
147#[derive(Clone, derive_new::new)]
148pub struct DeferralOutputFiller<F: VmField> {
149 count_chip: Arc<DeferralCircuitCountChip>,
150 poseidon2_chip: Arc<DeferralPoseidon2Chip<F>>,
151 bitwise_lookup_chip: SharedBitwiseOperationLookupChip<RV32_CELL_BITS>,
152 address_bits: usize,
153}
154
155impl<F, RA> PreflightExecutor<F, RA> for DeferralOutputExecutor
156where
157 F: PrimeField32,
158 for<'buf> RA: RecordArena<'buf, DeferralOutputLayout, DeferralOutputRecordMut<'buf>>,
159{
160 fn get_opcode_name(&self, _opcode: usize) -> String {
161 format!("{:?}", DeferralOpcode::OUTPUT)
162 }
163
164 fn execute(
165 &self,
166 state: VmStateMut<F, TracingMemory, RA>,
167 instruction: &Instruction<F>,
168 ) -> Result<(), ExecutionError> {
169 let Instruction { a, b, c, d, e, .. } = instruction;
170 debug_assert_eq!(d.as_canonical_u32(), RV32_REGISTER_AS);
171 debug_assert_eq!(e.as_canonical_u32(), RV32_MEMORY_AS);
172
173 let rd_ptr = a.as_canonical_u32();
174 let rs_ptr = b.as_canonical_u32();
175 let deferral_idx = c.as_canonical_u32();
176
177 let read_ptr = read_rv32_register(state.memory.data(), rs_ptr);
179 let output_key_chunks: [[u8; DEFAULT_BLOCK_SIZE]; OUTPUT_TOTAL_MEMORY_OPS] = from_fn(|i| {
180 memory_read(
181 state.memory.data(),
182 RV32_MEMORY_AS,
183 read_ptr + (i * DEFAULT_BLOCK_SIZE) as u32,
184 )
185 });
186 let output_key: [u8; OUTPUT_TOTAL_BYTES] = join_memory_ops(output_key_chunks);
187 let (output_commit, output_len) = split_output(output_key);
188
189 let output_len_val = u64::from_le_bytes(output_len) as usize;
190 let num_rows = output_len_val / DIGEST_SIZE + 1;
191 debug_assert!(output_len_val.is_multiple_of(DIGEST_SIZE));
192
193 let record = state
195 .ctx
196 .alloc(DeferralOutputLayout::new(DeferralOutputMetadata {
197 num_rows,
198 }));
199
200 record.header.from_pc = *state.pc;
201 record.header.from_timestamp = state.memory.timestamp();
202 record.header.rd_ptr = rd_ptr;
203 record.header.rs_ptr = rs_ptr;
204 record.header.deferral_idx = deferral_idx;
205 record.header.num_rows = num_rows as u32;
206
207 record.header.rd_val = tracing_read(
208 state.memory,
209 RV32_REGISTER_AS,
210 rd_ptr,
211 &mut record.header.rd_aux.prev_timestamp,
212 );
213 record.header.rs_val = tracing_read(
214 state.memory,
215 RV32_REGISTER_AS,
216 rs_ptr,
217 &mut record.header.rs_aux.prev_timestamp,
218 );
219
220 let input_ptr = u32::from_le_bytes(record.header.rs_val);
221 let output_ptr = u32::from_le_bytes(record.header.rd_val);
222 for chunk_idx in 0..OUTPUT_TOTAL_MEMORY_OPS {
223 tracing_read::<DEFAULT_BLOCK_SIZE>(
224 state.memory,
225 RV32_MEMORY_AS,
226 input_ptr + (chunk_idx * DEFAULT_BLOCK_SIZE) as u32,
227 &mut record.header.output_commit_and_len_aux[chunk_idx].prev_timestamp,
228 );
229 }
230
231 let output_raw =
232 state.streams.deferrals[deferral_idx as usize].get_output(&output_commit.to_vec());
233 debug_assert_eq!(output_raw.len(), output_len_val);
234
235 for (row_idx, output_chunk) in output_raw.chunks_exact(DIGEST_SIZE).enumerate() {
236 let row_output_ptr = output_ptr + (row_idx * DIGEST_SIZE) as u32;
237 for chunk_idx in 0..DIGEST_MEMORY_OPS {
238 let aux_idx = row_idx * DIGEST_MEMORY_OPS + chunk_idx;
239 tracing_write(
240 state.memory,
241 RV32_MEMORY_AS,
242 row_output_ptr + (chunk_idx * DEFAULT_BLOCK_SIZE) as u32,
243 memory_op_chunk(output_chunk, chunk_idx),
244 &mut record.write_aux[aux_idx].prev_timestamp,
245 &mut record.write_aux[aux_idx].prev_data,
246 );
247 }
248 record.write_bytes[row_idx * DIGEST_SIZE..(row_idx + 1) * DIGEST_SIZE]
249 .copy_from_slice(output_chunk);
250 }
251
252 *state.pc = state.pc.wrapping_add(DEFAULT_PC_STEP);
253 Ok(())
254 }
255}
256
257impl<F> TraceFiller<F> for DeferralOutputFiller<F>
258where
259 F: VmField,
260{
261 fn fill_trace(
262 &self,
263 mem_helper: &MemoryAuxColsFactory<F>,
264 trace_matrix: &mut RowMajorMatrix<F>,
265 rows_used: usize,
266 ) {
267 if rows_used == 0 {
268 return;
269 }
270
271 let width = trace_matrix.width;
272 debug_assert_eq!(width, DeferralOutputCols::<u8>::width());
273
274 let mut trace = &mut trace_matrix.values[..width * rows_used];
275
276 while !trace.is_empty() {
277 let header: &DeferralOutputRecordHeader =
281 unsafe { get_record_from_slice(&mut trace, ()) };
282 let num_rows = header.num_rows as usize;
283 let output_len = (num_rows - 1) * DIGEST_SIZE;
284 let (mut section_chunk, rest) = trace.split_at_mut(width * num_rows);
285
286 let (header, write_bytes, write_aux) = {
288 let record: DeferralOutputRecordMut = unsafe {
292 get_record_from_slice(
293 &mut section_chunk,
294 DeferralOutputLayout::new(DeferralOutputMetadata { num_rows }),
295 )
296 };
297 (
298 record.header.clone(),
299 record.write_bytes.to_vec(),
300 record.write_aux.to_vec(),
301 )
302 };
303
304 let mut initial_sponge_input = [F::ZERO; DIGEST_SIZE];
306 initial_sponge_input[0] = F::from_u32(header.deferral_idx);
307 initial_sponge_input[1] = F::from_usize(output_len);
308
309 let mut current_poseidon2_res = [F::ZERO; DIGEST_SIZE];
310 self.count_chip.add_count(header.deferral_idx);
311
312 let output_len_bytes = u32::try_from(output_len)
313 .expect("deferral output length should fit a u32")
314 .to_le_bytes();
315 let output_len_f = output_len_bytes.map(F::from_u8);
316
317 for (row_idx, row) in section_chunk.chunks_exact_mut(width).enumerate() {
318 let cols: &mut DeferralOutputCols<F> = row.borrow_mut();
319
320 cols.is_valid = F::ONE;
321 cols.is_first = F::from_bool(row_idx == 0);
322 cols.is_last = F::from_bool(row_idx + 1 == num_rows);
323 cols.section_idx = F::from_usize(row_idx);
324
325 cols.from_state.pc = F::from_u32(header.from_pc);
326 cols.from_state.timestamp = F::from_u32(header.from_timestamp);
327 cols.rd_ptr = F::from_u32(header.rd_ptr);
328 cols.rs_ptr = F::from_u32(header.rs_ptr);
329 cols.deferral_idx = F::from_u32(header.deferral_idx);
330
331 cols.rd_val = header.rd_val.map(F::from_u8);
332 cols.rs_val = header.rs_val.map(F::from_u8);
333
334 if row_idx == 0 {
335 debug_assert!(RV32_CELL_BITS * RV32_REGISTER_NUM_LIMBS >= self.address_bits);
336 let limb_shift_bits =
337 RV32_CELL_BITS * RV32_REGISTER_NUM_LIMBS - self.address_bits;
338
339 self.bitwise_lookup_chip.request_range(
340 (header.rd_val[RV32_REGISTER_NUM_LIMBS - 1] as u32) << limb_shift_bits,
341 (header.rs_val[RV32_REGISTER_NUM_LIMBS - 1] as u32) << limb_shift_bits,
342 );
343 self.bitwise_lookup_chip.request_range(
344 (output_len_bytes[RV32_REGISTER_NUM_LIMBS - 1] as u32) << limb_shift_bits,
345 0,
346 );
347
348 mem_helper.fill(
349 header.rd_aux.prev_timestamp,
350 header.from_timestamp,
351 cols.rd_aux.as_mut(),
352 );
353 mem_helper.fill(
354 header.rs_aux.prev_timestamp,
355 header.from_timestamp + 1,
356 cols.rs_aux.as_mut(),
357 );
358 for chunk_idx in 0..OUTPUT_TOTAL_MEMORY_OPS {
359 mem_helper.fill(
360 header.output_commit_and_len_aux[chunk_idx].prev_timestamp,
361 header.from_timestamp + 2 + chunk_idx as u32,
362 cols.output_commit_and_len_aux[chunk_idx].as_mut(),
363 );
364 }
365 } else {
366 mem_helper.fill_zero(cols.rd_aux.as_mut());
367 mem_helper.fill_zero(cols.rs_aux.as_mut());
368 for chunk_aux in &mut cols.output_commit_and_len_aux {
369 mem_helper.fill_zero(chunk_aux.as_mut());
370 }
371 for aux in &mut cols.output_commit_lt_aux {
377 CanonicityTraceGen::clear_aux(aux);
378 }
379 }
380
381 cols.output_len = output_len_f;
382 if row_idx == 0 {
383 cols.sponge_inputs = initial_sponge_input;
384 current_poseidon2_res = self.poseidon2_chip.perm_and_record(
385 &cols.sponge_inputs,
386 &[F::ZERO; DIGEST_SIZE],
387 row_idx + 1 == num_rows,
388 );
389 for chunk_aux in &mut cols.write_bytes_aux {
390 mem_helper.fill_zero(chunk_aux.as_mut());
391 }
392 } else {
393 let output_chunk =
394 &write_bytes[(row_idx - 1) * DIGEST_SIZE..row_idx * DIGEST_SIZE];
395 for bytes in output_chunk.chunks_exact(2) {
396 self.bitwise_lookup_chip
397 .request_range(bytes[0] as u32, bytes[1] as u32);
398 }
399 cols.sponge_inputs = from_fn(|i| F::from_u8(output_chunk[i]));
400 current_poseidon2_res = self.poseidon2_chip.perm_and_record(
401 &cols.sponge_inputs,
402 ¤t_poseidon2_res,
403 row_idx + 1 == num_rows,
404 );
405 for chunk_idx in 0..DIGEST_MEMORY_OPS {
406 let aux_idx = (row_idx - 1) * DIGEST_MEMORY_OPS + chunk_idx;
407 cols.write_bytes_aux[chunk_idx]
408 .set_prev_data(write_aux[aux_idx].prev_data.map(F::from_u8));
409 mem_helper.fill(
410 write_aux[aux_idx].prev_timestamp,
411 header.from_timestamp
412 + 2
413 + OUTPUT_TOTAL_MEMORY_OPS as u32
414 + aux_idx as u32,
415 cols.write_bytes_aux[chunk_idx].as_mut(),
416 );
417 }
418 }
419 cols.poseidon2_res = current_poseidon2_res;
420 }
421
422 let output_commit = f_commit_to_bytes(¤t_poseidon2_res).map(F::from_u8);
423 for row in section_chunk.chunks_exact_mut(width) {
424 let cols: &mut DeferralOutputCols<F> = row.borrow_mut();
425 cols.output_commit = output_commit;
426 }
427 let cols: &mut DeferralOutputCols<F> = section_chunk[..width].borrow_mut();
428 let output_commit_rcs = output_commit
429 .chunks_exact(F_NUM_BYTES)
430 .zip(cols.output_commit_lt_aux.iter_mut())
431 .map(|(bytes, aux)| {
432 let x_le = from_fn(|i| bytes[i]);
433 CanonicityTraceGen::generate_subrow(&x_le, aux)
434 })
435 .collect_vec();
436 for rc_pair in output_commit_rcs.chunks_exact(2) {
437 self.bitwise_lookup_chip
438 .request_range(rc_pair[0], rc_pair[1]);
439 }
440
441 trace = rest;
442 }
443
444 trace_matrix.values[width * rows_used..].fill(F::ZERO);
445 }
446}