openvm_deferral_circuit/output/
trace.rs

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    // Heap pointers and auxiliary records
72    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    // Output commit and length read auxiliary record
78    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        // SAFETY:
90        // - Caller guarantees through the layout that self has sufficient length for all splits
91        let (header_buf, rest) =
92            unsafe { self.split_at_mut_unchecked(size_of::<DeferralOutputRecordHeader>()) };
93
94        // SAFETY:
95        // - The layout guarantees rest has sufficient length for write data
96        // - There are DIGEST_SIZE bytes written per row
97        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        // SAFETY:
102        // - Valid mutable slice from the previous split
103        // - Middle slice is properly aligned for MemoryWriteBytesAuxRecord via align_to_mut
104        // - Subslice operation [..layout.metadata.num_rows] validates sufficient capacity
105        // - Layout calculation ensures space for alignment padding plus required aux records
106        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        // Do a non-tracing read to get the output_len and compute num_rows
178        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        // We now have the layout and can write the record
194        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            // SAFETY:
278            // - Executor writes a valid record to the start of trace
279            // - Header is at the start of the record
280            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            // Copy write data out first; row filling overwrites the record bytes in-place.
287            let (header, write_bytes, write_aux) = {
288                // SAFETY:
289                // - The section contains exactly one DeferralOutputRecord
290                // - Layout is reconstructed from the record header
291                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            // Initial sponge input is [deferral_idx, output_len, 0, ...].
305            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                    // The canonicity aux columns are only populated on the first row (the
372                    // canonicity range check is gated by `is_first`). On non-first rows the
373                    // preflight record may have left non-zero bytes in these columns, so clear
374                    // them to satisfy the unconditional `assert_bool` constraints in the
375                    // canonicity sub-AIR.
376                    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                        &current_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(&current_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}