openvm_rv32im_circuit/adapters/
rdwrite.rs

1use std::borrow::{Borrow, BorrowMut};
2
3use openvm_circuit::{
4    arch::{
5        get_record_from_slice, AdapterAirContext, AdapterTraceExecutor, AdapterTraceFiller,
6        BasicAdapterInterface, ExecutionBridge, ExecutionState, ImmInstruction, VmAdapterAir,
7    },
8    system::memory::{
9        offline_checker::{MemoryBridge, MemoryWriteAuxCols, MemoryWriteBytesAuxRecord},
10        online::TracingMemory,
11        MemoryAddress, MemoryAuxColsFactory,
12    },
13};
14use openvm_circuit_primitives::{
15    utils::not, AlignedBytesBorrow, ColumnsAir, StructReflection, StructReflectionHelper,
16};
17use openvm_circuit_primitives_derive::AlignedBorrow;
18use openvm_instructions::{
19    instruction::Instruction, program::DEFAULT_PC_STEP, riscv::RV32_REGISTER_AS,
20};
21use openvm_stark_backend::{
22    interaction::InteractionBuilder,
23    p3_air::{AirBuilder, BaseAir},
24    p3_field::{Field, PrimeCharacteristicRing, PrimeField32},
25};
26
27use super::RV32_REGISTER_NUM_LIMBS;
28use crate::adapters::tracing_write;
29
30#[repr(C)]
31#[derive(Debug, Clone, AlignedBorrow, StructReflection)]
32pub struct Rv32RdWriteAdapterCols<T> {
33    pub from_state: ExecutionState<T>,
34    pub rd_ptr: T,
35    pub rd_aux_cols: MemoryWriteAuxCols<T, RV32_REGISTER_NUM_LIMBS>,
36}
37
38#[repr(C)]
39#[derive(Debug, Clone, AlignedBorrow, StructReflection)]
40pub struct Rv32CondRdWriteAdapterCols<T> {
41    pub inner: Rv32RdWriteAdapterCols<T>,
42    pub needs_write: T,
43}
44
45/// This adapter doesn't read anything, and writes to \[a:4\]_d, where d == 1
46#[derive(Clone, Copy, Debug, derive_new::new, ColumnsAir)]
47#[columns_via(Rv32RdWriteAdapterCols<u8>)]
48pub struct Rv32RdWriteAdapterAir {
49    pub(super) memory_bridge: MemoryBridge,
50    pub(super) execution_bridge: ExecutionBridge,
51}
52
53/// This adapter doesn't read anything, and **maybe** writes to \[a:4\]_d, where d == 1
54#[derive(Clone, Copy, Debug, derive_new::new, ColumnsAir)]
55#[columns_via(Rv32CondRdWriteAdapterCols<u8>)]
56pub struct Rv32CondRdWriteAdapterAir {
57    inner: Rv32RdWriteAdapterAir,
58}
59
60impl<F: Field> BaseAir<F> for Rv32RdWriteAdapterAir {
61    fn width(&self) -> usize {
62        Rv32RdWriteAdapterCols::<F>::width()
63    }
64}
65
66impl<F: Field> BaseAir<F> for Rv32CondRdWriteAdapterAir {
67    fn width(&self) -> usize {
68        Rv32CondRdWriteAdapterCols::<F>::width()
69    }
70}
71
72impl Rv32RdWriteAdapterAir {
73    /// If `needs_write` is provided:
74    /// - Only writes if `needs_write`.
75    /// - Sets operand `f = needs_write` in the instruction.
76    /// - Does not put any other constraints on `needs_write`
77    ///
78    /// Otherwise:
79    /// - Writes if `ctx.instruction.is_valid`.
80    /// - Sets operand `f` to default value of `0` in the instruction.
81    #[allow(clippy::type_complexity)]
82    fn conditional_eval<AB: InteractionBuilder>(
83        &self,
84        builder: &mut AB,
85        local_cols: &Rv32RdWriteAdapterCols<AB::Var>,
86        ctx: AdapterAirContext<
87            AB::Expr,
88            BasicAdapterInterface<
89                AB::Expr,
90                ImmInstruction<AB::Expr>,
91                0,
92                1,
93                0,
94                RV32_REGISTER_NUM_LIMBS,
95            >,
96        >,
97        needs_write: Option<AB::Expr>,
98    ) {
99        let timestamp: AB::Var = local_cols.from_state.timestamp;
100        let timestamp_delta = 1;
101        let (write_count, f) = if let Some(needs_write) = needs_write {
102            (needs_write.clone(), needs_write)
103        } else {
104            (ctx.instruction.is_valid.clone(), AB::Expr::ZERO)
105        };
106        self.memory_bridge
107            .write(
108                MemoryAddress::new(AB::F::from_u32(RV32_REGISTER_AS), local_cols.rd_ptr),
109                ctx.writes[0].clone(),
110                timestamp,
111                &local_cols.rd_aux_cols,
112            )
113            .eval(builder, write_count);
114
115        let to_pc = ctx
116            .to_pc
117            .unwrap_or(local_cols.from_state.pc + AB::F::from_u32(DEFAULT_PC_STEP));
118        // regardless of `needs_write`, must always execute instruction when `is_valid`.
119        self.execution_bridge
120            .execute(
121                ctx.instruction.opcode,
122                [
123                    local_cols.rd_ptr.into(),
124                    AB::Expr::ZERO,
125                    ctx.instruction.immediate,
126                    AB::Expr::from_u32(RV32_REGISTER_AS),
127                    AB::Expr::ZERO,
128                    f,
129                ],
130                local_cols.from_state,
131                ExecutionState {
132                    pc: to_pc,
133                    timestamp: timestamp + AB::F::from_usize(timestamp_delta),
134                },
135            )
136            .eval(builder, ctx.instruction.is_valid);
137    }
138}
139
140impl<AB: InteractionBuilder> VmAdapterAir<AB> for Rv32RdWriteAdapterAir {
141    type Interface =
142        BasicAdapterInterface<AB::Expr, ImmInstruction<AB::Expr>, 0, 1, 0, RV32_REGISTER_NUM_LIMBS>;
143
144    fn eval(
145        &self,
146        builder: &mut AB,
147        local: &[AB::Var],
148        ctx: AdapterAirContext<AB::Expr, Self::Interface>,
149    ) {
150        let local_cols: &Rv32RdWriteAdapterCols<AB::Var> = (*local).borrow();
151        self.conditional_eval(builder, local_cols, ctx, None);
152    }
153
154    fn get_from_pc(&self, local: &[AB::Var]) -> AB::Var {
155        let cols: &Rv32RdWriteAdapterCols<_> = local.borrow();
156        cols.from_state.pc
157    }
158}
159
160impl<AB: InteractionBuilder> VmAdapterAir<AB> for Rv32CondRdWriteAdapterAir {
161    type Interface =
162        BasicAdapterInterface<AB::Expr, ImmInstruction<AB::Expr>, 0, 1, 0, RV32_REGISTER_NUM_LIMBS>;
163
164    fn eval(
165        &self,
166        builder: &mut AB,
167        local: &[AB::Var],
168        ctx: AdapterAirContext<AB::Expr, Self::Interface>,
169    ) {
170        let local_cols: &Rv32CondRdWriteAdapterCols<AB::Var> = (*local).borrow();
171
172        builder.assert_bool(local_cols.needs_write);
173        builder
174            .when::<AB::Expr>(not(ctx.instruction.is_valid.clone()))
175            .assert_zero(local_cols.needs_write);
176
177        self.inner.conditional_eval(
178            builder,
179            &local_cols.inner,
180            ctx,
181            Some(local_cols.needs_write.into()),
182        );
183    }
184
185    fn get_from_pc(&self, local: &[AB::Var]) -> AB::Var {
186        let cols: &Rv32CondRdWriteAdapterCols<_> = local.borrow();
187        cols.inner.from_state.pc
188    }
189}
190
191/// This adapter doesn't read anything, and writes to \[a:4\]_d, where d == 1
192#[repr(C)]
193#[derive(AlignedBytesBorrow, Debug, Clone)]
194pub struct Rv32RdWriteAdapterRecord {
195    pub from_pc: u32,
196    pub from_timestamp: u32,
197
198    // Will use u32::MAX to indicate no write
199    pub rd_ptr: u32,
200    pub rd_aux_record: MemoryWriteBytesAuxRecord<RV32_REGISTER_NUM_LIMBS>,
201}
202
203#[derive(Clone, Copy, derive_new::new)]
204pub struct Rv32RdWriteAdapterExecutor;
205
206#[derive(Clone, Copy, derive_new::new)]
207pub struct Rv32RdWriteAdapterFiller;
208
209impl<F> AdapterTraceExecutor<F> for Rv32RdWriteAdapterExecutor
210where
211    F: PrimeField32,
212{
213    const WIDTH: usize = size_of::<Rv32RdWriteAdapterCols<u8>>();
214    type ReadData = ();
215    type WriteData = [u8; RV32_REGISTER_NUM_LIMBS];
216    type RecordMut<'a> = &'a mut Rv32RdWriteAdapterRecord;
217
218    #[inline(always)]
219    fn start(pc: u32, memory: &TracingMemory, record: &mut Self::RecordMut<'_>) {
220        record.from_pc = pc;
221        record.from_timestamp = memory.timestamp;
222    }
223
224    #[inline(always)]
225    fn read(
226        &self,
227        _memory: &mut TracingMemory,
228        _instruction: &Instruction<F>,
229        _record: &mut Self::RecordMut<'_>,
230    ) -> Self::ReadData {
231        // Rv32RdWriteAdapter doesn't read anything
232    }
233
234    #[inline(always)]
235    fn write(
236        &self,
237        memory: &mut TracingMemory,
238        instruction: &Instruction<F>,
239        data: Self::WriteData,
240        record: &mut Self::RecordMut<'_>,
241    ) {
242        let &Instruction { a, d, .. } = instruction;
243
244        debug_assert_eq!(d.as_canonical_u32(), RV32_REGISTER_AS);
245
246        record.rd_ptr = a.as_canonical_u32();
247        tracing_write(
248            memory,
249            RV32_REGISTER_AS,
250            record.rd_ptr,
251            data,
252            &mut record.rd_aux_record.prev_timestamp,
253            &mut record.rd_aux_record.prev_data,
254        );
255    }
256}
257
258impl<F: PrimeField32> AdapterTraceFiller<F> for Rv32RdWriteAdapterFiller {
259    const WIDTH: usize = size_of::<Rv32RdWriteAdapterCols<u8>>();
260
261    #[inline(always)]
262    fn fill_trace_row(&self, mem_helper: &MemoryAuxColsFactory<F>, mut adapter_row: &mut [F]) {
263        // SAFETY:
264        // - caller ensures `adapter_row` contains a valid record representation that was previously
265        //   written by the executor
266        // - get_record_from_slice correctly interprets the bytes as Rv32RdWriteAdapterRecord
267        let record: &Rv32RdWriteAdapterRecord =
268            unsafe { get_record_from_slice(&mut adapter_row, ()) };
269        let adapter_row: &mut Rv32RdWriteAdapterCols<F> = adapter_row.borrow_mut();
270
271        adapter_row
272            .rd_aux_cols
273            .set_prev_data(record.rd_aux_record.prev_data.map(F::from_u8));
274        mem_helper.fill(
275            record.rd_aux_record.prev_timestamp,
276            record.from_timestamp,
277            adapter_row.rd_aux_cols.as_mut(),
278        );
279        adapter_row.rd_ptr = F::from_u32(record.rd_ptr);
280        adapter_row.from_state.timestamp = F::from_u32(record.from_timestamp);
281        adapter_row.from_state.pc = F::from_u32(record.from_pc);
282    }
283}
284
285/// This adapter doesn't read anything, and **maybe** writes to \[a:4\]_d, where d == 1
286#[derive(Clone, Copy, derive_new::new)]
287pub struct Rv32CondRdWriteAdapterExecutor {
288    inner: Rv32RdWriteAdapterExecutor,
289}
290
291#[derive(Clone, Copy, derive_new::new)]
292pub struct Rv32CondRdWriteAdapterFiller {
293    inner: Rv32RdWriteAdapterFiller,
294}
295
296impl<F> AdapterTraceExecutor<F> for Rv32CondRdWriteAdapterExecutor
297where
298    F: PrimeField32,
299{
300    const WIDTH: usize = size_of::<Rv32CondRdWriteAdapterCols<u8>>();
301    type ReadData = ();
302    type WriteData = [u8; RV32_REGISTER_NUM_LIMBS];
303    type RecordMut<'a> = &'a mut Rv32RdWriteAdapterRecord;
304
305    #[inline(always)]
306    fn start(pc: u32, memory: &TracingMemory, record: &mut Self::RecordMut<'_>) {
307        record.from_pc = pc;
308        record.from_timestamp = memory.timestamp;
309    }
310
311    #[inline(always)]
312    fn read(
313        &self,
314        memory: &mut TracingMemory,
315        instruction: &Instruction<F>,
316        record: &mut Self::RecordMut<'_>,
317    ) -> Self::ReadData {
318        <Rv32RdWriteAdapterExecutor as AdapterTraceExecutor<F>>::read(
319            &self.inner,
320            memory,
321            instruction,
322            record,
323        )
324    }
325
326    #[inline(always)]
327    fn write(
328        &self,
329        memory: &mut TracingMemory,
330        instruction: &Instruction<F>,
331        data: Self::WriteData,
332        record: &mut Self::RecordMut<'_>,
333    ) {
334        let Instruction { f: enabled, .. } = instruction;
335
336        if enabled.is_one() {
337            <Rv32RdWriteAdapterExecutor as AdapterTraceExecutor<F>>::write(
338                &self.inner,
339                memory,
340                instruction,
341                data,
342                record,
343            );
344        } else {
345            memory.increment_timestamp();
346            record.rd_ptr = u32::MAX;
347        }
348    }
349}
350
351impl<F: PrimeField32> AdapterTraceFiller<F> for Rv32CondRdWriteAdapterFiller {
352    const WIDTH: usize = size_of::<Rv32CondRdWriteAdapterCols<u8>>();
353
354    #[inline(always)]
355    fn fill_trace_row(&self, mem_helper: &MemoryAuxColsFactory<F>, mut adapter_row: &mut [F]) {
356        // SAFETY:
357        // - caller ensures `adapter_row` contains a valid record representation that was previously
358        //   written by the executor
359        // - get_record_from_slice correctly interprets the bytes as Rv32RdWriteAdapterRecord
360        let record: &Rv32RdWriteAdapterRecord =
361            unsafe { get_record_from_slice(&mut adapter_row, ()) };
362        let adapter_cols: &mut Rv32CondRdWriteAdapterCols<F> = adapter_row.borrow_mut();
363
364        adapter_cols.needs_write = F::from_bool(record.rd_ptr != u32::MAX);
365
366        if record.rd_ptr != u32::MAX {
367            // SAFETY:
368            // - adapter_row has sufficient length for the split
369            // - size_of::<Rv32RdWriteAdapterCols<u8>>() is the correct split point
370            unsafe {
371                self.inner.fill_trace_row(
372                    mem_helper,
373                    adapter_row
374                        .split_at_mut_unchecked(size_of::<Rv32RdWriteAdapterCols<u8>>())
375                        .0,
376                )
377            };
378        } else {
379            adapter_cols.inner.rd_ptr = F::ZERO;
380            mem_helper.fill_zero(adapter_cols.inner.rd_aux_cols.as_mut());
381            adapter_cols.inner.from_state.timestamp = F::from_u32(record.from_timestamp);
382            adapter_cols.inner.from_state.pc = F::from_u32(record.from_pc);
383        }
384    }
385}