openvm_bigint_circuit/
base_alu.rs

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        // Compiler will expand this loop.
256        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        // Compiler will expand this loop.
269        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        // Compiler will expand this loop.
282        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}