openvm_rv32im_circuit/adapters/
loadstore.rs

1use std::{
2    borrow::{Borrow, BorrowMut},
3    marker::PhantomData,
4};
5
6use openvm_circuit::{
7    arch::{
8        get_record_from_slice, AdapterAirContext, AdapterTraceExecutor, AdapterTraceFiller,
9        ExecutionBridge, ExecutionState, VmAdapterAir, VmAdapterInterface,
10    },
11    system::memory::{
12        offline_checker::{
13            MemoryBaseAuxCols, MemoryBridge, MemoryReadAuxCols, MemoryReadAuxRecord,
14            MemoryWriteAuxCols,
15        },
16        online::TracingMemory,
17        MemoryAddress, MemoryAuxColsFactory,
18    },
19};
20use openvm_circuit_primitives::{
21    utils::{not, select},
22    var_range::{SharedVariableRangeCheckerChip, VariableRangeCheckerBus},
23    AlignedBytesBorrow, ColumnsAir, StructReflection, StructReflectionHelper,
24};
25use openvm_circuit_primitives_derive::AlignedBorrow;
26use openvm_instructions::{
27    instruction::Instruction,
28    program::DEFAULT_PC_STEP,
29    riscv::{RV32_IMM_AS, RV32_MEMORY_AS, RV32_REGISTER_AS},
30    LocalOpcode, DEFERRAL_AS,
31};
32use openvm_rv32im_transpiler::Rv32LoadStoreOpcode::{self, *};
33use openvm_stark_backend::{
34    interaction::InteractionBuilder,
35    p3_air::{AirBuilder, BaseAir},
36    p3_field::{Field, PrimeCharacteristicRing, PrimeField32},
37};
38
39use super::RV32_REGISTER_NUM_LIMBS;
40use crate::adapters::{memory_read, timed_write, tracing_read, RV32_CELL_BITS};
41
42/// LoadStore Adapter handles all memory and register operations, so it must be aware
43/// of the instruction type, specifically whether it is a load or store
44/// LoadStore Adapter handles 4 byte aligned lw, sw instructions,
45///                           2 byte aligned lh, lhu, sh instructions and
46///                           1 byte aligned lb, lbu, sb instructions
47/// This adapter always batch reads/writes 4 bytes,
48/// thus it needs to shift left the memory pointer by some amount in case of not 4 byte aligned
49/// intermediate pointers
50pub struct LoadStoreInstruction<T> {
51    /// is_valid is constrained to be bool
52    pub is_valid: T,
53    /// Absolute opcode number
54    pub opcode: T,
55    /// is_load is constrained to be bool, and can only be 1 if is_valid is 1
56    pub is_load: T,
57
58    /// Keeping two separate shift amounts is needed for getting the read_ptr/write_ptr with degree
59    /// 2 load_shift_amount will be the shift amount if load and 0 if store
60    pub load_shift_amount: T,
61    /// store_shift_amount will be 0 if load and the shift amount if store
62    pub store_shift_amount: T,
63}
64
65pub struct Rv32LoadStoreAdapterAirInterface<AB: InteractionBuilder>(PhantomData<AB>);
66
67/// Using AB::Var for prev_data and AB::Expr for read_data
68impl<AB: InteractionBuilder> VmAdapterInterface<AB::Expr> for Rv32LoadStoreAdapterAirInterface<AB> {
69    type Reads = (
70        [AB::Var; RV32_REGISTER_NUM_LIMBS],
71        [AB::Expr; RV32_REGISTER_NUM_LIMBS],
72    );
73    type Writes = [[AB::Expr; RV32_REGISTER_NUM_LIMBS]; 1];
74    type ProcessedInstruction = LoadStoreInstruction<AB::Expr>;
75}
76
77#[repr(C)]
78#[derive(Debug, Clone, AlignedBorrow, StructReflection)]
79pub struct Rv32LoadStoreAdapterCols<T> {
80    pub from_state: ExecutionState<T>,
81    pub rs1_ptr: T,
82    pub rs1_data: [T; RV32_REGISTER_NUM_LIMBS],
83    pub rs1_aux_cols: MemoryReadAuxCols<T>,
84
85    /// Will write to rd when Load and read from rs2 when Store
86    pub rd_rs2_ptr: T,
87    pub read_data_aux: MemoryReadAuxCols<T>,
88    pub imm: T,
89    pub imm_sign: T,
90    /// mem_ptr is the intermediate memory pointer limbs, needed to check the correct addition
91    pub mem_ptr_limbs: [T; 2],
92    pub mem_as: T,
93    /// prev_data will be provided by the core chip to make a complete MemoryWriteAuxCols
94    pub write_base_aux: MemoryBaseAuxCols<T>,
95    /// Only writes if `needs_write`.
96    /// If the instruction is a Load:
97    /// - Sets `needs_write` to 0 iff `rd == x0`
98    ///
99    /// Otherwise:
100    /// - Sets `needs_write` to 1
101    pub needs_write: T,
102}
103
104#[derive(Clone, Copy, Debug, derive_new::new, ColumnsAir)]
105#[columns_via(Rv32LoadStoreAdapterCols<u8>)]
106pub struct Rv32LoadStoreAdapterAir {
107    pub(super) memory_bridge: MemoryBridge,
108    pub(super) execution_bridge: ExecutionBridge,
109    pub range_bus: VariableRangeCheckerBus,
110    pointer_max_bits: usize,
111}
112
113impl<F: Field> BaseAir<F> for Rv32LoadStoreAdapterAir {
114    fn width(&self) -> usize {
115        Rv32LoadStoreAdapterCols::<F>::width()
116    }
117}
118
119impl<AB: InteractionBuilder> VmAdapterAir<AB> for Rv32LoadStoreAdapterAir {
120    type Interface = Rv32LoadStoreAdapterAirInterface<AB>;
121
122    fn eval(
123        &self,
124        builder: &mut AB,
125        local: &[AB::Var],
126        ctx: AdapterAirContext<AB::Expr, Self::Interface>,
127    ) {
128        let local_cols: &Rv32LoadStoreAdapterCols<AB::Var> = local.borrow();
129
130        let timestamp: AB::Var = local_cols.from_state.timestamp;
131        let mut timestamp_delta: usize = 0;
132        let mut timestamp_pp = || {
133            timestamp_delta += 1;
134            timestamp + AB::Expr::from_usize(timestamp_delta - 1)
135        };
136
137        let is_load = ctx.instruction.is_load;
138        let is_valid = ctx.instruction.is_valid;
139        let load_shift_amount = ctx.instruction.load_shift_amount;
140        let store_shift_amount = ctx.instruction.store_shift_amount;
141        let shift_amount = load_shift_amount.clone() + store_shift_amount.clone();
142
143        let write_count = local_cols.needs_write;
144
145        // This constraint ensures that the memory write only occurs when `is_valid == 1`.
146        builder.assert_bool(write_count);
147        builder.when(write_count).assert_one(is_valid.clone());
148
149        // Constrain that if `is_valid == 1` and `write_count == 0`, then `is_load == 1` and
150        // `rd_rs2_ptr == x0`
151        builder
152            .when(is_valid.clone() - write_count)
153            .assert_one(is_load.clone());
154        builder
155            .when(is_valid.clone() - write_count)
156            .assert_zero(local_cols.rd_rs2_ptr);
157
158        // read rs1
159        self.memory_bridge
160            .read(
161                MemoryAddress::new(AB::F::from_u32(RV32_REGISTER_AS), local_cols.rs1_ptr),
162                local_cols.rs1_data,
163                timestamp_pp(),
164                &local_cols.rs1_aux_cols,
165            )
166            .eval(builder, is_valid.clone());
167
168        // constrain mem_ptr = rs1 + imm as a u32 addition with 2 limbs
169        let limbs_01 =
170            local_cols.rs1_data[0] + local_cols.rs1_data[1] * AB::F::from_u32(1 << RV32_CELL_BITS);
171        let limbs_23 =
172            local_cols.rs1_data[2] + local_cols.rs1_data[3] * AB::F::from_u32(1 << RV32_CELL_BITS);
173
174        let inv = AB::F::from_u32(1 << (RV32_CELL_BITS * 2)).inverse();
175        let carry = (limbs_01 + local_cols.imm - local_cols.mem_ptr_limbs[0]) * inv;
176
177        builder.when(is_valid.clone()).assert_bool(carry.clone());
178
179        builder
180            .when(is_valid.clone())
181            .assert_bool(local_cols.imm_sign);
182        let imm_extend_limb =
183            local_cols.imm_sign * AB::F::from_u32((1 << (RV32_CELL_BITS * 2)) - 1);
184        let carry = (limbs_23 + imm_extend_limb + carry - local_cols.mem_ptr_limbs[1]) * inv;
185        builder.when(is_valid.clone()).assert_bool(carry.clone());
186
187        // preventing mem_ptr overflow
188        self.range_bus
189            .range_check(
190                // (limb[0] - shift_amount) / 4 < 2^14 => limb[0] - shift_amount < 2^16
191                (local_cols.mem_ptr_limbs[0] - shift_amount) * AB::F::from_u32(4).inverse(),
192                RV32_CELL_BITS * 2 - 2,
193            )
194            .eval(builder, is_valid.clone());
195        self.range_bus
196            .range_check(
197                local_cols.mem_ptr_limbs[1],
198                self.pointer_max_bits - RV32_CELL_BITS * 2,
199            )
200            .eval(builder, is_valid.clone());
201
202        let mem_ptr = local_cols.mem_ptr_limbs[0]
203            + local_cols.mem_ptr_limbs[1] * AB::F::from_u32(1 << (RV32_CELL_BITS * 2));
204
205        let is_store = is_valid.clone() - is_load.clone();
206        // constrain mem_as to be in {0, 1, 2} if the instruction is a load,
207        // and in {2, 3, 4} if the instruction is a store
208        builder.assert_tern(local_cols.mem_as - is_store * AB::Expr::TWO);
209        builder
210            .when(not::<AB::Expr>(is_valid.clone()))
211            .assert_zero(local_cols.mem_as);
212
213        // read_as is [local_cols.mem_as] for loads and 1 for stores
214        let read_as = select::<AB::Expr>(
215            is_load.clone(),
216            local_cols.mem_as,
217            AB::F::from_u32(RV32_REGISTER_AS),
218        );
219
220        // read_ptr is mem_ptr for loads and rd_rs2_ptr for stores
221        // Note: shift_amount is expected to have degree 2, thus we can't put it in the select
222        // clause       since the resulting read_ptr/write_ptr's degree will be 3 which is
223        // too high.       Instead, the solution without using additional columns is to get
224        // two different shift amounts from core chip
225        let read_ptr = select::<AB::Expr>(is_load.clone(), mem_ptr.clone(), local_cols.rd_rs2_ptr)
226            - load_shift_amount;
227
228        self.memory_bridge
229            .read(
230                MemoryAddress::new(read_as, read_ptr),
231                ctx.reads.1,
232                timestamp_pp(),
233                &local_cols.read_data_aux,
234            )
235            .eval(builder, is_valid.clone());
236
237        let write_aux_cols = MemoryWriteAuxCols::from_base(local_cols.write_base_aux, ctx.reads.0);
238
239        // write_as is 1 for loads and [local_cols.mem_as] for stores
240        let write_as = select::<AB::Expr>(
241            is_load.clone(),
242            AB::F::from_u32(RV32_REGISTER_AS),
243            local_cols.mem_as,
244        );
245
246        // write_ptr is rd_rs2_ptr for loads and mem_ptr for stores
247        let write_ptr = select::<AB::Expr>(is_load.clone(), local_cols.rd_rs2_ptr, mem_ptr.clone())
248            - store_shift_amount;
249
250        self.memory_bridge
251            .write(
252                MemoryAddress::new(write_as, write_ptr),
253                ctx.writes[0].clone(),
254                timestamp_pp(),
255                &write_aux_cols,
256            )
257            .eval(builder, write_count);
258
259        let to_pc = ctx
260            .to_pc
261            .unwrap_or(local_cols.from_state.pc + AB::F::from_u32(DEFAULT_PC_STEP));
262        self.execution_bridge
263            .execute(
264                ctx.instruction.opcode,
265                [
266                    local_cols.rd_rs2_ptr.into(),
267                    local_cols.rs1_ptr.into(),
268                    local_cols.imm.into(),
269                    AB::Expr::from_u32(RV32_REGISTER_AS),
270                    local_cols.mem_as.into(),
271                    local_cols.needs_write.into(),
272                    local_cols.imm_sign.into(),
273                ],
274                local_cols.from_state,
275                ExecutionState {
276                    pc: to_pc,
277                    timestamp: timestamp + AB::F::from_usize(timestamp_delta),
278                },
279            )
280            .eval(builder, is_valid);
281    }
282
283    fn get_from_pc(&self, local: &[AB::Var]) -> AB::Var {
284        let local_cols: &Rv32LoadStoreAdapterCols<AB::Var> = local.borrow();
285        local_cols.from_state.pc
286    }
287}
288
289#[repr(C)]
290#[derive(AlignedBytesBorrow, Debug)]
291pub struct Rv32LoadStoreAdapterRecord {
292    pub from_pc: u32,
293    pub from_timestamp: u32,
294
295    pub rs1_ptr: u32,
296    pub rs1_val: u32,
297    pub rs1_aux_record: MemoryReadAuxRecord,
298
299    pub rd_rs2_ptr: u32,
300    pub read_data_aux: MemoryReadAuxRecord,
301    pub imm: u16,
302    pub imm_sign: bool,
303
304    pub mem_as: u8,
305
306    pub write_prev_timestamp: u32,
307}
308
309/// This chip reads rs1 and gets a intermediate memory pointer address with rs1 + imm.
310/// In case of Loads, reads from the shifted intermediate pointer and writes to rd.
311/// In case of Stores, reads from rs2 and writes to the shifted intermediate pointer.
312#[derive(Clone, Copy, derive_new::new)]
313pub struct Rv32LoadStoreAdapterExecutor {
314    pointer_max_bits: usize,
315}
316
317#[derive(derive_new::new)]
318pub struct Rv32LoadStoreAdapterFiller {
319    pointer_max_bits: usize,
320    pub range_checker_chip: SharedVariableRangeCheckerChip,
321}
322
323impl<F> AdapterTraceExecutor<F> for Rv32LoadStoreAdapterExecutor
324where
325    F: PrimeField32,
326{
327    const WIDTH: usize = size_of::<Rv32LoadStoreAdapterCols<u8>>();
328    type ReadData = (
329        (
330            [u32; RV32_REGISTER_NUM_LIMBS],
331            [u8; RV32_REGISTER_NUM_LIMBS],
332        ),
333        u8,
334    );
335    type WriteData = [u32; RV32_REGISTER_NUM_LIMBS];
336    type RecordMut<'a> = &'a mut Rv32LoadStoreAdapterRecord;
337
338    #[inline(always)]
339    fn start(pc: u32, memory: &TracingMemory, record: &mut Self::RecordMut<'_>) {
340        record.from_pc = pc;
341        record.from_timestamp = memory.timestamp;
342    }
343
344    #[inline(always)]
345    fn read(
346        &self,
347        memory: &mut TracingMemory,
348        instruction: &Instruction<F>,
349        record: &mut Self::RecordMut<'_>,
350    ) -> Self::ReadData {
351        let &Instruction {
352            opcode,
353            a,
354            b,
355            c,
356            d,
357            e,
358            g,
359            ..
360        } = instruction;
361
362        debug_assert_eq!(d.as_canonical_u32(), RV32_REGISTER_AS);
363
364        let local_opcode = Rv32LoadStoreOpcode::from_usize(
365            opcode.local_opcode_idx(Rv32LoadStoreOpcode::CLASS_OFFSET),
366        );
367
368        record.rs1_ptr = b.as_canonical_u32();
369        record.rs1_val = u32::from_le_bytes(tracing_read(
370            memory,
371            RV32_REGISTER_AS,
372            record.rs1_ptr,
373            &mut record.rs1_aux_record.prev_timestamp,
374        ));
375
376        record.imm = c.as_canonical_u32() as u16;
377        record.imm_sign = g.is_one();
378        let imm_extended = record.imm as u32 + record.imm_sign as u32 * 0xffff0000;
379
380        let ptr_val = record.rs1_val.wrapping_add(imm_extended);
381        let shift_amount = ptr_val & 3;
382        let ptr_val = ptr_val - shift_amount;
383
384        assert!(
385            ptr_val < (1 << self.pointer_max_bits),
386            "ptr_val: {ptr_val} = rs1_val: {} + imm_extended: {imm_extended} >= 2 ** {}",
387            record.rs1_val,
388            self.pointer_max_bits
389        );
390
391        // prev_data: We need to keep values of some cells to keep them unchanged when writing to
392        // those cells
393        let (read_data, prev_data) = match local_opcode {
394            LOADW | LOADB | LOADH | LOADBU | LOADHU => {
395                debug_assert_eq!(e, F::from_u32(RV32_MEMORY_AS));
396                record.mem_as = RV32_MEMORY_AS as u8;
397                let read_data = tracing_read(
398                    memory,
399                    RV32_MEMORY_AS,
400                    ptr_val,
401                    &mut record.read_data_aux.prev_timestamp,
402                );
403                let prev_data = memory_read(memory.data(), RV32_REGISTER_AS, a.as_canonical_u32())
404                    .map(u32::from);
405                (read_data, prev_data)
406            }
407            STOREW | STOREH | STOREB => {
408                let e = e.as_canonical_u32();
409                debug_assert_ne!(e, RV32_IMM_AS);
410                debug_assert_ne!(e, RV32_REGISTER_AS);
411                debug_assert_ne!(e, DEFERRAL_AS);
412                record.mem_as = e as u8;
413                let read_data = tracing_read(
414                    memory,
415                    RV32_REGISTER_AS,
416                    a.as_canonical_u32(),
417                    &mut record.read_data_aux.prev_timestamp,
418                );
419                let prev_data = memory_read(memory.data(), e, ptr_val).map(u32::from);
420                (read_data, prev_data)
421            }
422        };
423
424        ((prev_data, read_data), shift_amount as u8)
425    }
426
427    #[inline(always)]
428    fn write(
429        &self,
430        memory: &mut TracingMemory,
431        instruction: &Instruction<F>,
432        data: Self::WriteData,
433        record: &mut Self::RecordMut<'_>,
434    ) {
435        let &Instruction {
436            opcode,
437            a,
438            d,
439            e,
440            f: enabled,
441            ..
442        } = instruction;
443
444        debug_assert_eq!(d.as_canonical_u32(), RV32_REGISTER_AS);
445        debug_assert_ne!(e.as_canonical_u32(), RV32_IMM_AS);
446        debug_assert_ne!(e.as_canonical_u32(), RV32_REGISTER_AS);
447        debug_assert_ne!(e.as_canonical_u32(), DEFERRAL_AS);
448
449        let local_opcode = Rv32LoadStoreOpcode::from_usize(
450            opcode.local_opcode_idx(Rv32LoadStoreOpcode::CLASS_OFFSET),
451        );
452
453        if enabled != F::ZERO {
454            record.rd_rs2_ptr = a.as_canonical_u32();
455
456            record.write_prev_timestamp = match local_opcode {
457                STOREW | STOREH | STOREB => {
458                    let imm_extended = record.imm as u32 + record.imm_sign as u32 * 0xffff0000;
459                    let ptr = record.rs1_val.wrapping_add(imm_extended) & !3;
460
461                    timed_write(memory, record.mem_as as u32, ptr, data.map(|x| x as u8)).0
462                }
463                LOADW | LOADB | LOADH | LOADBU | LOADHU => {
464                    timed_write(
465                        memory,
466                        RV32_REGISTER_AS,
467                        record.rd_rs2_ptr,
468                        data.map(|x| x as u8),
469                    )
470                    .0
471                }
472            };
473        } else {
474            record.rd_rs2_ptr = u32::MAX;
475            memory.increment_timestamp();
476        };
477    }
478}
479
480impl<F: PrimeField32> AdapterTraceFiller<F> for Rv32LoadStoreAdapterFiller {
481    const WIDTH: usize = size_of::<Rv32LoadStoreAdapterCols<u8>>();
482
483    #[inline(always)]
484    fn fill_trace_row(&self, mem_helper: &MemoryAuxColsFactory<F>, mut adapter_row: &mut [F]) {
485        debug_assert!(self.range_checker_chip.range_max_bits() >= 15);
486
487        // SAFETY:
488        // - caller ensures `adapter_row` contains a valid record representation that was previously
489        //   written by the executor
490        // - get_record_from_slice correctly interprets the bytes as Rv32LoadStoreAdapterRecord
491        let record: &Rv32LoadStoreAdapterRecord =
492            unsafe { get_record_from_slice(&mut adapter_row, ()) };
493        let adapter_row: &mut Rv32LoadStoreAdapterCols<F> = adapter_row.borrow_mut();
494
495        let needs_write = record.rd_rs2_ptr != u32::MAX;
496        // Writing in reverse order
497        adapter_row.needs_write = F::from_bool(needs_write);
498
499        if needs_write {
500            mem_helper.fill(
501                record.write_prev_timestamp,
502                record.from_timestamp + 2,
503                &mut adapter_row.write_base_aux,
504            );
505        } else {
506            mem_helper.fill_zero(&mut adapter_row.write_base_aux);
507        }
508
509        adapter_row.mem_as = F::from_u8(record.mem_as);
510        let ptr = record
511            .rs1_val
512            .wrapping_add(record.imm as u32 + record.imm_sign as u32 * 0xffff0000);
513
514        let ptr_limbs = [ptr & 0xffff, ptr >> 16];
515        self.range_checker_chip
516            .add_count(ptr_limbs[0] >> 2, RV32_CELL_BITS * 2 - 2);
517        self.range_checker_chip
518            .add_count(ptr_limbs[1], self.pointer_max_bits - 16);
519        adapter_row.mem_ptr_limbs = ptr_limbs.map(F::from_u32);
520
521        adapter_row.imm_sign = F::from_bool(record.imm_sign);
522        adapter_row.imm = F::from_u16(record.imm);
523
524        mem_helper.fill(
525            record.read_data_aux.prev_timestamp,
526            record.from_timestamp + 1,
527            adapter_row.read_data_aux.as_mut(),
528        );
529        adapter_row.rd_rs2_ptr = if record.rd_rs2_ptr != u32::MAX {
530            F::from_u32(record.rd_rs2_ptr)
531        } else {
532            F::ZERO
533        };
534
535        mem_helper.fill(
536            record.rs1_aux_record.prev_timestamp,
537            record.from_timestamp,
538            adapter_row.rs1_aux_cols.as_mut(),
539        );
540
541        adapter_row.rs1_data = record.rs1_val.to_le_bytes().map(F::from_u8);
542        adapter_row.rs1_ptr = F::from_u32(record.rs1_ptr);
543
544        adapter_row.from_state.timestamp = F::from_u32(record.from_timestamp);
545        adapter_row.from_state.pc = F::from_u32(record.from_pc);
546    }
547}