1use std::{
2 borrow::{Borrow, BorrowMut},
3 mem::size_of,
4};
5
6use openvm_bigint_transpiler::Rv32BaseAlu256Opcode;
7use openvm_circuit::{arch::*, system::memory::online::GuestMemory};
8use openvm_circuit_primitives_derive::AlignedBytesBorrow;
9use openvm_instructions::{
10 instruction::Instruction,
11 program::DEFAULT_PC_STEP,
12 riscv::{RV32_MEMORY_AS, RV32_REGISTER_AS},
13 LocalOpcode,
14};
15use openvm_rv32im_circuit::BaseAluExecutor;
16use openvm_rv32im_transpiler::BaseAluOpcode;
17use openvm_stark_backend::p3_field::PrimeField32;
18
19use crate::{
20 common::{bytes_to_u64_array, read_int256, u64_array_to_bytes, write_int256},
21 AluAdapterExecutor, Rv32BaseAlu256Executor, INT256_NUM_LIMBS,
22};
23
24impl Rv32BaseAlu256Executor {
25 pub fn new(adapter: AluAdapterExecutor, offset: usize) -> Self {
26 Self(BaseAluExecutor::new(adapter, offset))
27 }
28}
29
30#[derive(AlignedBytesBorrow)]
31struct BaseAluPreCompute {
32 a: u8,
33 b: u8,
34 c: u8,
35}
36
37macro_rules! dispatch {
38 ($execute_impl:ident, $local_opcode:ident) => {
39 Ok(match $local_opcode {
40 BaseAluOpcode::ADD => $execute_impl::<_, _, AddOp>,
41 BaseAluOpcode::SUB => $execute_impl::<_, _, SubOp>,
42 BaseAluOpcode::XOR => $execute_impl::<_, _, XorOp>,
43 BaseAluOpcode::OR => $execute_impl::<_, _, OrOp>,
44 BaseAluOpcode::AND => $execute_impl::<_, _, AndOp>,
45 })
46 };
47}
48
49impl<F: PrimeField32> InterpreterExecutor<F> for Rv32BaseAlu256Executor {
50 fn pre_compute_size(&self) -> usize {
51 size_of::<BaseAluPreCompute>()
52 }
53
54 #[cfg(not(feature = "tco"))]
55 fn pre_compute<Ctx>(
56 &self,
57 pc: u32,
58 inst: &Instruction<F>,
59 data: &mut [u8],
60 ) -> Result<ExecuteFunc<F, Ctx>, StaticProgramError>
61 where
62 Ctx: ExecutionCtxTrait,
63 {
64 let data: &mut BaseAluPreCompute = data.borrow_mut();
65 let local_opcode = self.pre_compute_impl(pc, inst, data)?;
66
67 dispatch!(execute_e1_handler, local_opcode)
68 }
69
70 #[cfg(feature = "tco")]
71 fn handler<Ctx>(
72 &self,
73 pc: u32,
74 inst: &Instruction<F>,
75 data: &mut [u8],
76 ) -> Result<Handler<F, Ctx>, StaticProgramError>
77 where
78 Ctx: ExecutionCtxTrait,
79 {
80 let data: &mut BaseAluPreCompute = data.borrow_mut();
81 let local_opcode = self.pre_compute_impl(pc, inst, data)?;
82
83 dispatch!(execute_e1_handler, local_opcode)
84 }
85}
86
87#[cfg(feature = "aot")]
88impl<F: PrimeField32> AotExecutor<F> for Rv32BaseAlu256Executor {}
89
90impl<F: PrimeField32> InterpreterMeteredExecutor<F> for Rv32BaseAlu256Executor {
91 fn metered_pre_compute_size(&self) -> usize {
92 size_of::<E2PreCompute<BaseAluPreCompute>>()
93 }
94
95 #[cfg(not(feature = "tco"))]
96 fn metered_pre_compute<Ctx>(
97 &self,
98 chip_idx: usize,
99 pc: u32,
100 inst: &Instruction<F>,
101 data: &mut [u8],
102 ) -> Result<ExecuteFunc<F, Ctx>, StaticProgramError>
103 where
104 Ctx: MeteredExecutionCtxTrait,
105 {
106 let data: &mut E2PreCompute<BaseAluPreCompute> = data.borrow_mut();
107 data.chip_idx = chip_idx as u32;
108 let local_opcode = self.pre_compute_impl(pc, inst, &mut data.data)?;
109
110 dispatch!(execute_e2_handler, local_opcode)
111 }
112
113 #[cfg(feature = "tco")]
114 fn metered_handler<Ctx>(
115 &self,
116 chip_idx: usize,
117 pc: u32,
118 inst: &Instruction<F>,
119 data: &mut [u8],
120 ) -> Result<Handler<F, Ctx>, StaticProgramError>
121 where
122 Ctx: MeteredExecutionCtxTrait,
123 {
124 let data: &mut E2PreCompute<BaseAluPreCompute> = data.borrow_mut();
125 data.chip_idx = chip_idx as u32;
126 let local_opcode = self.pre_compute_impl(pc, inst, &mut data.data)?;
127
128 dispatch!(execute_e2_handler, local_opcode)
129 }
130}
131#[cfg(feature = "aot")]
132impl<F: PrimeField32> AotMeteredExecutor<F> for Rv32BaseAlu256Executor {}
133
134#[inline(always)]
135unsafe fn execute_e12_impl<F: PrimeField32, CTX: ExecutionCtxTrait, OP: AluOp>(
136 pre_compute: &BaseAluPreCompute,
137 exec_state: &mut VmExecState<F, GuestMemory, CTX>,
138) {
139 let rs1_ptr = exec_state.vm_read::<u8, 4>(RV32_REGISTER_AS, pre_compute.b as u32);
140 let rs2_ptr = exec_state.vm_read::<u8, 4>(RV32_REGISTER_AS, pre_compute.c as u32);
141 let rd_ptr = exec_state.vm_read::<u8, 4>(RV32_REGISTER_AS, pre_compute.a as u32);
142 let rs1 = read_int256(exec_state, RV32_MEMORY_AS, u32::from_le_bytes(rs1_ptr));
143 let rs2 = read_int256(exec_state, RV32_MEMORY_AS, u32::from_le_bytes(rs2_ptr));
144 let rd = <OP as AluOp>::compute(rs1, rs2);
145 write_int256(exec_state, RV32_MEMORY_AS, u32::from_le_bytes(rd_ptr), &rd);
146 let pc = exec_state.pc();
147 exec_state.set_pc(pc.wrapping_add(DEFAULT_PC_STEP));
148}
149
150#[create_handler]
151#[inline(always)]
152unsafe fn execute_e1_impl<F: PrimeField32, CTX: ExecutionCtxTrait, OP: AluOp>(
153 pre_compute: *const u8,
154 exec_state: &mut VmExecState<F, GuestMemory, CTX>,
155) {
156 let pre_compute: &BaseAluPreCompute =
157 std::slice::from_raw_parts(pre_compute, size_of::<BaseAluPreCompute>()).borrow();
158 execute_e12_impl::<F, CTX, OP>(pre_compute, exec_state);
159}
160
161#[create_handler]
162#[inline(always)]
163unsafe fn execute_e2_impl<F: PrimeField32, CTX: MeteredExecutionCtxTrait, OP: AluOp>(
164 pre_compute: *const u8,
165 exec_state: &mut VmExecState<F, GuestMemory, CTX>,
166) {
167 let pre_compute: &E2PreCompute<BaseAluPreCompute> =
168 std::slice::from_raw_parts(pre_compute, size_of::<E2PreCompute<BaseAluPreCompute>>())
169 .borrow();
170 exec_state
171 .ctx
172 .on_height_change(pre_compute.chip_idx as usize, 1);
173 execute_e12_impl::<F, CTX, OP>(&pre_compute.data, exec_state);
174}
175
176impl Rv32BaseAlu256Executor {
177 fn pre_compute_impl<F: PrimeField32>(
178 &self,
179 pc: u32,
180 inst: &Instruction<F>,
181 data: &mut BaseAluPreCompute,
182 ) -> Result<BaseAluOpcode, StaticProgramError> {
183 let Instruction {
184 opcode,
185 a,
186 b,
187 c,
188 d,
189 e,
190 ..
191 } = inst;
192 let e_u32 = e.as_canonical_u32();
193 if d.as_canonical_u32() != RV32_REGISTER_AS || e_u32 != RV32_MEMORY_AS {
194 return Err(StaticProgramError::InvalidInstruction(pc));
195 }
196 *data = BaseAluPreCompute {
197 a: a.as_canonical_u32() as u8,
198 b: b.as_canonical_u32() as u8,
199 c: c.as_canonical_u32() as u8,
200 };
201 let local_opcode =
202 BaseAluOpcode::from_usize(opcode.local_opcode_idx(Rv32BaseAlu256Opcode::CLASS_OFFSET));
203 Ok(local_opcode)
204 }
205}
206
207trait AluOp {
208 fn compute(rs1: [u8; INT256_NUM_LIMBS], rs2: [u8; INT256_NUM_LIMBS]) -> [u8; INT256_NUM_LIMBS];
209}
210struct AddOp;
211struct SubOp;
212struct XorOp;
213struct OrOp;
214struct AndOp;
215impl AluOp for AddOp {
216 #[inline(always)]
217 fn compute(rs1: [u8; INT256_NUM_LIMBS], rs2: [u8; INT256_NUM_LIMBS]) -> [u8; INT256_NUM_LIMBS] {
218 let rs1_u64: [u64; 4] = bytes_to_u64_array(rs1);
219 let rs2_u64: [u64; 4] = bytes_to_u64_array(rs2);
220 let mut rd_u64 = [0u64; 4];
221 let (res, mut carry) = rs1_u64[0].overflowing_add(rs2_u64[0]);
222 rd_u64[0] = res;
223 for i in 1..4 {
224 let (res1, c1) = rs1_u64[i].overflowing_add(rs2_u64[i]);
225 let (res2, c2) = res1.overflowing_add(carry as u64);
226 carry = c1 || c2;
227 rd_u64[i] = res2;
228 }
229 u64_array_to_bytes(rd_u64)
230 }
231}
232impl AluOp for SubOp {
233 #[inline(always)]
234 fn compute(rs1: [u8; INT256_NUM_LIMBS], rs2: [u8; INT256_NUM_LIMBS]) -> [u8; INT256_NUM_LIMBS] {
235 let rs1_u64: [u64; 4] = bytes_to_u64_array(rs1);
236 let rs2_u64: [u64; 4] = bytes_to_u64_array(rs2);
237 let mut rd_u64 = [0u64; 4];
238 let (res, mut borrow) = rs1_u64[0].overflowing_sub(rs2_u64[0]);
239 rd_u64[0] = res;
240 for i in 1..4 {
241 let (res1, c1) = rs1_u64[i].overflowing_sub(rs2_u64[i]);
242 let (res2, c2) = res1.overflowing_sub(borrow as u64);
243 borrow = c1 || c2;
244 rd_u64[i] = res2;
245 }
246 u64_array_to_bytes(rd_u64)
247 }
248}
249impl AluOp for XorOp {
250 #[inline(always)]
251 fn compute(rs1: [u8; INT256_NUM_LIMBS], rs2: [u8; INT256_NUM_LIMBS]) -> [u8; INT256_NUM_LIMBS] {
252 let rs1_u64: [u64; 4] = bytes_to_u64_array(rs1);
253 let rs2_u64: [u64; 4] = bytes_to_u64_array(rs2);
254 let mut rd_u64 = [0u64; 4];
255 for i in 0..4 {
257 rd_u64[i] = rs1_u64[i] ^ rs2_u64[i];
258 }
259 u64_array_to_bytes(rd_u64)
260 }
261}
262impl AluOp for OrOp {
263 #[inline(always)]
264 fn compute(rs1: [u8; INT256_NUM_LIMBS], rs2: [u8; INT256_NUM_LIMBS]) -> [u8; INT256_NUM_LIMBS] {
265 let rs1_u64: [u64; 4] = bytes_to_u64_array(rs1);
266 let rs2_u64: [u64; 4] = bytes_to_u64_array(rs2);
267 let mut rd_u64 = [0u64; 4];
268 for i in 0..4 {
270 rd_u64[i] = rs1_u64[i] | rs2_u64[i];
271 }
272 u64_array_to_bytes(rd_u64)
273 }
274}
275impl AluOp for AndOp {
276 #[inline(always)]
277 fn compute(rs1: [u8; INT256_NUM_LIMBS], rs2: [u8; INT256_NUM_LIMBS]) -> [u8; INT256_NUM_LIMBS] {
278 let rs1_u64: [u64; 4] = bytes_to_u64_array(rs1);
279 let rs2_u64: [u64; 4] = bytes_to_u64_array(rs2);
280 let mut rd_u64 = [0u64; 4];
281 for i in 0..4 {
283 rd_u64[i] = rs1_u64[i] & rs2_u64[i];
284 }
285 u64_array_to_bytes(rd_u64)
286 }
287}