openvm_rv32im_circuit/adapters/
mul.rs

1use std::borrow::{Borrow, BorrowMut};
2
3use openvm_circuit::{
4    arch::{
5        get_record_from_slice, AdapterAirContext, AdapterTraceExecutor, AdapterTraceFiller,
6        BasicAdapterInterface, ExecutionBridge, ExecutionState, MinimalInstruction, VmAdapterAir,
7    },
8    system::memory::{
9        offline_checker::{
10            MemoryBridge, MemoryReadAuxCols, MemoryReadAuxRecord, MemoryWriteAuxCols,
11            MemoryWriteBytesAuxRecord,
12        },
13        online::TracingMemory,
14        MemoryAddress, MemoryAuxColsFactory,
15    },
16};
17use openvm_circuit_primitives::{
18    AlignedBytesBorrow, ColumnsAir, StructReflection, StructReflectionHelper,
19};
20use openvm_circuit_primitives_derive::AlignedBorrow;
21use openvm_instructions::{
22    instruction::Instruction, program::DEFAULT_PC_STEP, riscv::RV32_REGISTER_AS,
23};
24use openvm_stark_backend::{
25    interaction::InteractionBuilder,
26    p3_air::BaseAir,
27    p3_field::{Field, PrimeCharacteristicRing, PrimeField32},
28};
29
30use super::{tracing_write, RV32_REGISTER_NUM_LIMBS};
31use crate::adapters::tracing_read;
32
33#[repr(C)]
34#[derive(AlignedBorrow, StructReflection)]
35pub struct Rv32MultAdapterCols<T> {
36    pub from_state: ExecutionState<T>,
37    pub rd_ptr: T,
38    pub rs1_ptr: T,
39    pub rs2_ptr: T,
40    pub reads_aux: [MemoryReadAuxCols<T>; 2],
41    pub writes_aux: MemoryWriteAuxCols<T, RV32_REGISTER_NUM_LIMBS>,
42}
43
44/// Reads instructions of the form OP a, b, c, d where \[a:4\]_d = \[b:4\]_d op \[c:4\]_d.
45/// Operand d can only be 1, and there is no immediate support.
46#[derive(Clone, Copy, Debug, derive_new::new, ColumnsAir)]
47#[columns_via(Rv32MultAdapterCols<u8>)]
48pub struct Rv32MultAdapterAir {
49    pub(super) execution_bridge: ExecutionBridge,
50    pub(super) memory_bridge: MemoryBridge,
51}
52
53impl<F: Field> BaseAir<F> for Rv32MultAdapterAir {
54    fn width(&self) -> usize {
55        Rv32MultAdapterCols::<F>::width()
56    }
57}
58
59impl<AB: InteractionBuilder> VmAdapterAir<AB> for Rv32MultAdapterAir {
60    type Interface = BasicAdapterInterface<
61        AB::Expr,
62        MinimalInstruction<AB::Expr>,
63        2,
64        1,
65        RV32_REGISTER_NUM_LIMBS,
66        RV32_REGISTER_NUM_LIMBS,
67    >;
68
69    fn eval(
70        &self,
71        builder: &mut AB,
72        local: &[AB::Var],
73        ctx: AdapterAirContext<AB::Expr, Self::Interface>,
74    ) {
75        let local: &Rv32MultAdapterCols<_> = local.borrow();
76        let timestamp = local.from_state.timestamp;
77        let mut timestamp_delta: usize = 0;
78        let mut timestamp_pp = || {
79            timestamp_delta += 1;
80            timestamp + AB::F::from_usize(timestamp_delta - 1)
81        };
82
83        self.memory_bridge
84            .read(
85                MemoryAddress::new(AB::F::from_u32(RV32_REGISTER_AS), local.rs1_ptr),
86                ctx.reads[0].clone(),
87                timestamp_pp(),
88                &local.reads_aux[0],
89            )
90            .eval(builder, ctx.instruction.is_valid.clone());
91
92        self.memory_bridge
93            .read(
94                MemoryAddress::new(AB::F::from_u32(RV32_REGISTER_AS), local.rs2_ptr),
95                ctx.reads[1].clone(),
96                timestamp_pp(),
97                &local.reads_aux[1],
98            )
99            .eval(builder, ctx.instruction.is_valid.clone());
100
101        self.memory_bridge
102            .write(
103                MemoryAddress::new(AB::F::from_u32(RV32_REGISTER_AS), local.rd_ptr),
104                ctx.writes[0].clone(),
105                timestamp_pp(),
106                &local.writes_aux,
107            )
108            .eval(builder, ctx.instruction.is_valid.clone());
109
110        self.execution_bridge
111            .execute_and_increment_or_set_pc(
112                ctx.instruction.opcode,
113                [
114                    local.rd_ptr.into(),
115                    local.rs1_ptr.into(),
116                    local.rs2_ptr.into(),
117                    AB::Expr::from_u32(RV32_REGISTER_AS),
118                    AB::Expr::ZERO,
119                ],
120                local.from_state,
121                AB::F::from_usize(timestamp_delta),
122                (DEFAULT_PC_STEP, ctx.to_pc),
123            )
124            .eval(builder, ctx.instruction.is_valid);
125    }
126
127    fn get_from_pc(&self, local: &[AB::Var]) -> AB::Var {
128        let cols: &Rv32MultAdapterCols<_> = local.borrow();
129        cols.from_state.pc
130    }
131}
132
133#[repr(C)]
134#[derive(AlignedBytesBorrow, Debug)]
135pub struct Rv32MultAdapterRecord {
136    pub from_pc: u32,
137    pub from_timestamp: u32,
138
139    pub rd_ptr: u32,
140    pub rs1_ptr: u32,
141    pub rs2_ptr: u32,
142
143    pub reads_aux: [MemoryReadAuxRecord; 2],
144    pub writes_aux: MemoryWriteBytesAuxRecord<RV32_REGISTER_NUM_LIMBS>,
145}
146
147#[derive(Clone, Copy, derive_new::new)]
148pub struct Rv32MultAdapterExecutor;
149
150#[derive(Clone, Copy, derive_new::new)]
151pub struct Rv32MultAdapterFiller;
152
153impl<F> AdapterTraceExecutor<F> for Rv32MultAdapterExecutor
154where
155    F: PrimeField32,
156{
157    const WIDTH: usize = size_of::<Rv32MultAdapterCols<u8>>();
158    type ReadData = [[u8; RV32_REGISTER_NUM_LIMBS]; 2];
159    type WriteData = [[u8; RV32_REGISTER_NUM_LIMBS]; 1];
160    type RecordMut<'a> = &'a mut Rv32MultAdapterRecord;
161
162    #[inline(always)]
163    fn start(pc: u32, memory: &TracingMemory, record: &mut Self::RecordMut<'_>) {
164        record.from_pc = pc;
165        record.from_timestamp = memory.timestamp;
166    }
167
168    #[inline(always)]
169    fn read(
170        &self,
171        memory: &mut TracingMemory,
172        instruction: &Instruction<F>,
173        record: &mut Self::RecordMut<'_>,
174    ) -> Self::ReadData {
175        let &Instruction { b, c, d, .. } = instruction;
176
177        debug_assert_eq!(d.as_canonical_u32(), RV32_REGISTER_AS);
178
179        record.rs1_ptr = b.as_canonical_u32();
180        let rs1 = tracing_read(
181            memory,
182            RV32_REGISTER_AS,
183            b.as_canonical_u32(),
184            &mut record.reads_aux[0].prev_timestamp,
185        );
186        record.rs2_ptr = c.as_canonical_u32();
187        let rs2 = tracing_read(
188            memory,
189            RV32_REGISTER_AS,
190            c.as_canonical_u32(),
191            &mut record.reads_aux[1].prev_timestamp,
192        );
193
194        [rs1, rs2]
195    }
196
197    #[inline(always)]
198    fn write(
199        &self,
200        memory: &mut TracingMemory,
201        instruction: &Instruction<F>,
202        data: Self::WriteData,
203        record: &mut Self::RecordMut<'_>,
204    ) {
205        let &Instruction { a, d, .. } = instruction;
206
207        debug_assert_eq!(d.as_canonical_u32(), RV32_REGISTER_AS);
208
209        record.rd_ptr = a.as_canonical_u32();
210        tracing_write(
211            memory,
212            RV32_REGISTER_AS,
213            a.as_canonical_u32(),
214            data[0],
215            &mut record.writes_aux.prev_timestamp,
216            &mut record.writes_aux.prev_data,
217        )
218    }
219}
220
221impl<F: PrimeField32> AdapterTraceFiller<F> for Rv32MultAdapterFiller {
222    const WIDTH: usize = size_of::<Rv32MultAdapterCols<u8>>();
223
224    #[inline(always)]
225    fn fill_trace_row(&self, mem_helper: &MemoryAuxColsFactory<F>, mut adapter_row: &mut [F]) {
226        // SAFETY:
227        // - caller ensures `adapter_row` contains a valid record representation that was previously
228        //   written by the executor
229        // - get_record_from_slice correctly interprets the bytes as Rv32MultAdapterRecord
230        let record: &Rv32MultAdapterRecord = unsafe { get_record_from_slice(&mut adapter_row, ()) };
231        let adapter_row: &mut Rv32MultAdapterCols<F> = adapter_row.borrow_mut();
232
233        let timestamp = record.from_timestamp;
234
235        adapter_row
236            .writes_aux
237            .set_prev_data(record.writes_aux.prev_data.map(F::from_u8));
238        mem_helper.fill(
239            record.writes_aux.prev_timestamp,
240            timestamp + 2,
241            adapter_row.writes_aux.as_mut(),
242        );
243
244        mem_helper.fill(
245            record.reads_aux[1].prev_timestamp,
246            timestamp + 1,
247            adapter_row.reads_aux[1].as_mut(),
248        );
249
250        mem_helper.fill(
251            record.reads_aux[0].prev_timestamp,
252            timestamp,
253            adapter_row.reads_aux[0].as_mut(),
254        );
255
256        adapter_row.rs2_ptr = F::from_u32(record.rs2_ptr);
257        adapter_row.rs1_ptr = F::from_u32(record.rs1_ptr);
258        adapter_row.rd_ptr = F::from_u32(record.rd_ptr);
259
260        adapter_row.from_state.timestamp = F::from_u32(record.from_timestamp);
261        adapter_row.from_state.pc = F::from_u32(record.from_pc);
262    }
263}