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
42pub struct LoadStoreInstruction<T> {
51 pub is_valid: T,
53 pub opcode: T,
55 pub is_load: T,
57
58 pub load_shift_amount: T,
61 pub store_shift_amount: T,
63}
64
65pub struct Rv32LoadStoreAdapterAirInterface<AB: InteractionBuilder>(PhantomData<AB>);
66
67impl<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 pub rd_rs2_ptr: T,
87 pub read_data_aux: MemoryReadAuxCols<T>,
88 pub imm: T,
89 pub imm_sign: T,
90 pub mem_ptr_limbs: [T; 2],
92 pub mem_as: T,
93 pub write_base_aux: MemoryBaseAuxCols<T>,
95 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 builder.assert_bool(write_count);
147 builder.when(write_count).assert_one(is_valid.clone());
148
149 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 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 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 self.range_bus
189 .range_check(
190 (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 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 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 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 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 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#[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 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 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 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}