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#[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#[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 #[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 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#[repr(C)]
193#[derive(AlignedBytesBorrow, Debug, Clone)]
194pub struct Rv32RdWriteAdapterRecord {
195 pub from_pc: u32,
196 pub from_timestamp: u32,
197
198 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 }
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 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#[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 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 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}