1use std::{
2 array,
3 borrow::{Borrow, BorrowMut},
4};
5
6use openvm_circuit::{
7 arch::*,
8 system::memory::{online::TracingMemory, MemoryAuxColsFactory},
9};
10use openvm_circuit_primitives::{
11 bitwise_op_lookup::{BitwiseOperationLookupBus, SharedBitwiseOperationLookupChip},
12 range_tuple::{RangeTupleCheckerBus, SharedRangeTupleCheckerChip},
13 AlignedBytesBorrow, ColumnsAir, StructReflection, StructReflectionHelper,
14};
15use openvm_circuit_primitives_derive::AlignedBorrow;
16use openvm_instructions::{instruction::Instruction, program::DEFAULT_PC_STEP, LocalOpcode};
17use openvm_rv32im_transpiler::MulHOpcode;
18use openvm_stark_backend::{
19 interaction::InteractionBuilder,
20 p3_air::{AirBuilder, BaseAir},
21 p3_field::{Field, PrimeCharacteristicRing, PrimeField32},
22 BaseAirWithPublicValues,
23};
24use strum::IntoEnumIterator;
25
26#[repr(C)]
27#[derive(AlignedBorrow, StructReflection)]
28pub struct MulHCoreCols<T, const NUM_LIMBS: usize, const LIMB_BITS: usize> {
29 pub a: [T; NUM_LIMBS],
30 pub b: [T; NUM_LIMBS],
31 pub c: [T; NUM_LIMBS],
32
33 pub a_mul: [T; NUM_LIMBS],
34 pub b_ext: T,
35 pub c_ext: T,
36
37 pub opcode_mulh_flag: T,
38 pub opcode_mulhsu_flag: T,
39 pub opcode_mulhu_flag: T,
40}
41
42#[derive(Copy, Clone, Debug, derive_new::new, ColumnsAir)]
43#[columns_via(MulHCoreCols<u8, NUM_LIMBS, LIMB_BITS>)]
44pub struct MulHCoreAir<const NUM_LIMBS: usize, const LIMB_BITS: usize> {
45 pub bitwise_lookup_bus: BitwiseOperationLookupBus,
46 pub range_tuple_bus: RangeTupleCheckerBus<2>,
47}
48
49impl<F: Field, const NUM_LIMBS: usize, const LIMB_BITS: usize> BaseAir<F>
50 for MulHCoreAir<NUM_LIMBS, LIMB_BITS>
51{
52 fn width(&self) -> usize {
53 MulHCoreCols::<F, NUM_LIMBS, LIMB_BITS>::width()
54 }
55}
56impl<F: Field, const NUM_LIMBS: usize, const LIMB_BITS: usize> BaseAirWithPublicValues<F>
57 for MulHCoreAir<NUM_LIMBS, LIMB_BITS>
58{
59}
60
61impl<AB, I, const NUM_LIMBS: usize, const LIMB_BITS: usize> VmCoreAir<AB, I>
62 for MulHCoreAir<NUM_LIMBS, LIMB_BITS>
63where
64 AB: InteractionBuilder,
65 I: VmAdapterInterface<AB::Expr>,
66 I::Reads: From<[[AB::Expr; NUM_LIMBS]; 2]>,
67 I::Writes: From<[[AB::Expr; NUM_LIMBS]; 1]>,
68 I::ProcessedInstruction: From<MinimalInstruction<AB::Expr>>,
69{
70 fn eval(
71 &self,
72 builder: &mut AB,
73 local_core: &[AB::Var],
74 _from_pc: AB::Var,
75 ) -> AdapterAirContext<AB::Expr, I> {
76 let cols: &MulHCoreCols<_, NUM_LIMBS, LIMB_BITS> = local_core.borrow();
77 let flags = [
78 cols.opcode_mulh_flag,
79 cols.opcode_mulhsu_flag,
80 cols.opcode_mulhu_flag,
81 ];
82
83 let is_valid = flags.iter().fold(AB::Expr::ZERO, |acc, &flag| {
84 builder.assert_bool(flag);
85 acc + flag.into()
86 });
87 builder.assert_bool(is_valid.clone());
88
89 let b = &cols.b;
90 let c = &cols.c;
91 let carry_divide = AB::F::from_u32(1 << LIMB_BITS).inverse();
92
93 let a_mul = &cols.a_mul;
96 let mut carry_mul: [AB::Expr; NUM_LIMBS] = array::from_fn(|_| AB::Expr::ZERO);
97
98 for i in 0..NUM_LIMBS {
99 let expected_limb = if i == 0 {
100 AB::Expr::ZERO
101 } else {
102 carry_mul[i - 1].clone()
103 } + (0..=i).fold(AB::Expr::ZERO, |ac, k| ac + (b[k] * c[i - k]));
104 carry_mul[i] = AB::Expr::from(carry_divide) * (expected_limb - a_mul[i]);
105 }
106
107 for (a_mul, carry_mul) in a_mul.iter().zip(carry_mul.iter()) {
108 self.range_tuple_bus
109 .send(vec![(*a_mul).into(), carry_mul.clone()])
110 .eval(builder, is_valid.clone());
111 }
112
113 let a = &cols.a;
115 let mut carry: [AB::Expr; NUM_LIMBS] = array::from_fn(|_| AB::Expr::ZERO);
116
117 for j in 0..NUM_LIMBS {
118 let expected_limb = if j == 0 {
119 carry_mul[NUM_LIMBS - 1].clone()
120 } else {
121 carry[j - 1].clone()
122 } + ((j + 1)..NUM_LIMBS)
123 .fold(AB::Expr::ZERO, |acc, k| acc + (b[k] * c[NUM_LIMBS + j - k]))
124 + (0..(j + 1)).fold(AB::Expr::ZERO, |acc, k| {
125 acc + (b[k] * cols.c_ext) + (c[k] * cols.b_ext)
126 });
127 carry[j] = AB::Expr::from(carry_divide) * (expected_limb - a[j]);
128 }
129
130 for (a, carry) in a.iter().zip(carry.iter()) {
131 self.range_tuple_bus
132 .send(vec![(*a).into(), carry.clone()])
133 .eval(builder, is_valid.clone());
134 }
135
136 let sign_mask = AB::F::from_u32(1 << (LIMB_BITS - 1));
139 let ext_inv = AB::F::from_u32((1 << LIMB_BITS) - 1).inverse();
140 let b_sign = cols.b_ext * ext_inv;
141 let c_sign = cols.c_ext * ext_inv;
142
143 builder.assert_bool(b_sign.clone());
144 builder.assert_bool(c_sign.clone());
145 builder
146 .when(cols.opcode_mulhu_flag)
147 .assert_zero(b_sign.clone());
148 builder
149 .when(cols.opcode_mulhu_flag + cols.opcode_mulhsu_flag)
150 .assert_zero(c_sign.clone());
151
152 self.bitwise_lookup_bus
153 .send_range(
154 AB::Expr::from_u32(2) * (b[NUM_LIMBS - 1] - b_sign * sign_mask),
155 (cols.opcode_mulh_flag + AB::Expr::ONE) * (c[NUM_LIMBS - 1] - c_sign * sign_mask),
156 )
157 .eval(builder, cols.opcode_mulh_flag + cols.opcode_mulhsu_flag);
158
159 let expected_opcode = VmCoreAir::<AB, I>::expr_to_global_expr(
160 self,
161 flags.iter().zip(MulHOpcode::iter()).fold(
162 AB::Expr::ZERO,
163 |acc, (flag, local_opcode)| {
164 acc + (*flag).into() * AB::Expr::from_u8(local_opcode as u8)
165 },
166 ),
167 );
168
169 AdapterAirContext {
170 to_pc: None,
171 reads: [cols.b.map(Into::into), cols.c.map(Into::into)].into(),
172 writes: [cols.a.map(Into::into)].into(),
173 instruction: MinimalInstruction {
174 is_valid,
175 opcode: expected_opcode,
176 }
177 .into(),
178 }
179 }
180
181 fn start_offset(&self) -> usize {
182 MulHOpcode::CLASS_OFFSET
183 }
184}
185
186#[repr(C)]
187#[derive(AlignedBytesBorrow, Debug)]
188pub struct MulHCoreRecord<const NUM_LIMBS: usize, const LIMB_BITS: usize> {
189 pub b: [u8; NUM_LIMBS],
190 pub c: [u8; NUM_LIMBS],
191 pub local_opcode: u8,
192}
193
194#[derive(Clone, Copy, derive_new::new)]
195pub struct MulHExecutor<A, const NUM_LIMBS: usize, const LIMB_BITS: usize> {
196 adapter: A,
197 pub offset: usize,
198}
199
200#[derive(Clone)]
201pub struct MulHFiller<A, const NUM_LIMBS: usize, const LIMB_BITS: usize> {
202 adapter: A,
203 pub bitwise_lookup_chip: SharedBitwiseOperationLookupChip<LIMB_BITS>,
204 pub range_tuple_chip: SharedRangeTupleCheckerChip<2>,
205}
206
207impl<A, const NUM_LIMBS: usize, const LIMB_BITS: usize> MulHFiller<A, NUM_LIMBS, LIMB_BITS> {
208 pub fn new(
209 adapter: A,
210 bitwise_lookup_chip: SharedBitwiseOperationLookupChip<LIMB_BITS>,
211 range_tuple_chip: SharedRangeTupleCheckerChip<2>,
212 ) -> Self {
213 debug_assert!(
217 range_tuple_chip.sizes()[0] == 1 << LIMB_BITS,
218 "First element of RangeTupleChecker must have size {}",
219 1 << LIMB_BITS
220 );
221 debug_assert!(
222 range_tuple_chip.sizes()[1] >= (1 << LIMB_BITS) * 2 * NUM_LIMBS as u32,
223 "Second element of RangeTupleChecker must have size of at least {}",
224 (1 << LIMB_BITS) * 2 * NUM_LIMBS as u32
225 );
226
227 Self {
228 adapter,
229 bitwise_lookup_chip,
230 range_tuple_chip,
231 }
232 }
233}
234
235impl<F, A, RA, const NUM_LIMBS: usize, const LIMB_BITS: usize> PreflightExecutor<F, RA>
236 for MulHExecutor<A, NUM_LIMBS, LIMB_BITS>
237where
238 F: PrimeField32,
239 A: 'static
240 + AdapterTraceExecutor<
241 F,
242 ReadData: Into<[[u8; NUM_LIMBS]; 2]>,
243 WriteData: From<[[u8; NUM_LIMBS]; 1]>,
244 >,
245 for<'buf> RA: RecordArena<
246 'buf,
247 EmptyAdapterCoreLayout<F, A>,
248 (
249 A::RecordMut<'buf>,
250 &'buf mut MulHCoreRecord<NUM_LIMBS, LIMB_BITS>,
251 ),
252 >,
253{
254 fn get_opcode_name(&self, opcode: usize) -> String {
255 format!(
256 "{:?}",
257 MulHOpcode::from_usize(opcode - MulHOpcode::CLASS_OFFSET)
258 )
259 }
260
261 fn execute(
262 &self,
263 state: VmStateMut<F, TracingMemory, RA>,
264 instruction: &Instruction<F>,
265 ) -> Result<(), ExecutionError> {
266 let Instruction { opcode, .. } = instruction;
267
268 let (mut adapter_record, core_record) = state.ctx.alloc(EmptyAdapterCoreLayout::new());
269
270 A::start(*state.pc, state.memory, &mut adapter_record);
271
272 core_record.local_opcode = opcode.local_opcode_idx(MulHOpcode::CLASS_OFFSET) as u8;
273 let mulh_opcode = MulHOpcode::from_usize(core_record.local_opcode as usize);
274
275 [core_record.b, core_record.c] = self
276 .adapter
277 .read(state.memory, instruction, &mut adapter_record)
278 .into();
279
280 let (a, _, _, _, _) = run_mulh::<NUM_LIMBS, LIMB_BITS>(
281 mulh_opcode,
282 &core_record.b.map(u32::from),
283 &core_record.c.map(u32::from),
284 );
285
286 let a = a.map(|x| x as u8);
287 self.adapter
288 .write(state.memory, instruction, [a].into(), &mut adapter_record);
289
290 *state.pc = state.pc.wrapping_add(DEFAULT_PC_STEP);
291
292 Ok(())
293 }
294}
295
296impl<F, A, const NUM_LIMBS: usize, const LIMB_BITS: usize> TraceFiller<F>
297 for MulHFiller<A, NUM_LIMBS, LIMB_BITS>
298where
299 F: PrimeField32,
300 A: 'static + AdapterTraceFiller<F>,
301{
302 fn fill_trace_row(&self, mem_helper: &MemoryAuxColsFactory<F>, row_slice: &mut [F]) {
303 let (adapter_row, mut core_row) = unsafe { row_slice.split_at_mut_unchecked(A::WIDTH) };
306 self.adapter.fill_trace_row(mem_helper, adapter_row);
307 let record: &MulHCoreRecord<NUM_LIMBS, LIMB_BITS> =
310 unsafe { get_record_from_slice(&mut core_row, ()) };
311 let core_row: &mut MulHCoreCols<F, NUM_LIMBS, LIMB_BITS> = core_row.borrow_mut();
312
313 let opcode = MulHOpcode::from_usize(record.local_opcode as usize);
314 let (a, a_mul, carry, b_ext, c_ext) = run_mulh::<NUM_LIMBS, LIMB_BITS>(
315 opcode,
316 &record.b.map(u32::from),
317 &record.c.map(u32::from),
318 );
319
320 for i in 0..NUM_LIMBS {
321 self.range_tuple_chip.add_count(&[a_mul[i], carry[i]]);
322 self.range_tuple_chip
323 .add_count(&[a[i], carry[NUM_LIMBS + i]]);
324 }
325
326 if opcode != MulHOpcode::MULHU {
327 let b_sign_mask = if b_ext == 0 { 0 } else { 1 << (LIMB_BITS - 1) };
328 let c_sign_mask = if c_ext == 0 { 0 } else { 1 << (LIMB_BITS - 1) };
329 self.bitwise_lookup_chip.request_range(
330 (record.b[NUM_LIMBS - 1] as u32 - b_sign_mask) << 1,
331 (record.c[NUM_LIMBS - 1] as u32 - c_sign_mask)
332 << ((opcode == MulHOpcode::MULH) as u32),
333 );
334 }
335
336 core_row.opcode_mulhu_flag = F::from_bool(opcode == MulHOpcode::MULHU);
338 core_row.opcode_mulhsu_flag = F::from_bool(opcode == MulHOpcode::MULHSU);
339 core_row.opcode_mulh_flag = F::from_bool(opcode == MulHOpcode::MULH);
340 core_row.c_ext = F::from_u32(c_ext);
341 core_row.b_ext = F::from_u32(b_ext);
342 core_row.a_mul = a_mul.map(F::from_u32);
343 core_row.c = record.c.map(F::from_u8);
344 core_row.b = record.b.map(F::from_u8);
345 core_row.a = a.map(F::from_u32);
346 }
347}
348
349#[inline(always)]
351pub(super) fn run_mulh<const NUM_LIMBS: usize, const LIMB_BITS: usize>(
352 opcode: MulHOpcode,
353 x: &[u32; NUM_LIMBS],
354 y: &[u32; NUM_LIMBS],
355) -> ([u32; NUM_LIMBS], [u32; NUM_LIMBS], Vec<u32>, u32, u32) {
356 let mut mul = [0; NUM_LIMBS];
357 let mut carry = vec![0; 2 * NUM_LIMBS];
358 for i in 0..NUM_LIMBS {
359 if i > 0 {
360 mul[i] = carry[i - 1];
361 }
362 for j in 0..=i {
363 mul[i] += x[j] * y[i - j];
364 }
365 carry[i] = mul[i] >> LIMB_BITS;
366 mul[i] %= 1 << LIMB_BITS;
367 }
368
369 let x_ext = (x[NUM_LIMBS - 1] >> (LIMB_BITS - 1))
370 * if opcode == MulHOpcode::MULHU {
371 0
372 } else {
373 (1 << LIMB_BITS) - 1
374 };
375 let y_ext = (y[NUM_LIMBS - 1] >> (LIMB_BITS - 1))
376 * if opcode == MulHOpcode::MULH {
377 (1 << LIMB_BITS) - 1
378 } else {
379 0
380 };
381
382 let mut mulh = [0; NUM_LIMBS];
383 let mut x_prefix = 0;
384 let mut y_prefix = 0;
385
386 for i in 0..NUM_LIMBS {
387 x_prefix += x[i];
388 y_prefix += y[i];
389 mulh[i] = carry[NUM_LIMBS + i - 1] + x_prefix * y_ext + y_prefix * x_ext;
390 for j in (i + 1)..NUM_LIMBS {
391 mulh[i] += x[j] * y[NUM_LIMBS + i - j];
392 }
393 carry[NUM_LIMBS + i] = mulh[i] >> LIMB_BITS;
394 mulh[i] %= 1 << LIMB_BITS;
395 }
396
397 (mulh, mul, carry, x_ext, y_ext)
398}