1use acir::{
2 AcirField, BlackBoxFunc,
3 circuit::opcodes::FunctionInput,
4 native_types::{Witness, WitnessMap},
5};
6use acvm_blackbox_solver::BlackBoxFunctionSolver;
7
8use crate::pwg::{OpcodeResolutionError, input_to_value, insert_value};
9
10pub(super) fn multi_scalar_mul<F: AcirField>(
11 backend: &impl BlackBoxFunctionSolver<F>,
12 initial_witness: &mut WitnessMap<F>,
13 points: &[FunctionInput<F>],
14 scalars: &[FunctionInput<F>],
15 predicate: FunctionInput<F>,
16 outputs: (Witness, Witness),
17) -> Result<(), OpcodeResolutionError<F>> {
18 let (res_x, res_y) =
19 execute_multi_scalar_mul(backend, initial_witness, points, scalars, predicate)?;
20
21 insert_value(&outputs.0, res_x, initial_witness)?;
23 insert_value(&outputs.1, res_y, initial_witness)?;
24 Ok(())
25}
26
27pub(crate) fn execute_multi_scalar_mul<F: AcirField>(
28 backend: &impl BlackBoxFunctionSolver<F>,
29 initial_witness: &WitnessMap<F>,
30 points: &[FunctionInput<F>],
31 scalars: &[FunctionInput<F>],
32 predicate: FunctionInput<F>,
33) -> Result<(F, F), OpcodeResolutionError<F>> {
34 assert!(scalars.len().is_multiple_of(2), "Number of scalars must be even");
35 assert!(points.len().is_multiple_of(2), "Number of points must be a multiple of 2");
36 assert_eq!(
37 scalars.len() / 2,
38 points.len() / 2,
39 "Number of scalars must be the same as the number of points"
40 );
41
42 for point in points.chunks(2) {
43 if let [x, y] = *point {
44 check_all_or_nothing_pair(BlackBoxFunc::MultiScalarMul, "Coordinates", x, y)?;
45 }
46 }
47
48 for scalar in scalars.chunks(2) {
49 if let [lo, hi] = *scalar {
50 check_all_or_nothing_pair(BlackBoxFunc::MultiScalarMul, "Scalar limbs", lo, hi)?;
51 }
52 }
53
54 let points: Result<Vec<_>, _> =
55 points.iter().map(|input| input_to_value(initial_witness, *input)).collect();
56 let points: Vec<_> = points?.into_iter().collect();
57
58 let scalars: Result<Vec<_>, _> =
59 scalars.iter().map(|input| input_to_value(initial_witness, *input)).collect();
60
61 let predicate = input_to_value(initial_witness, predicate)?.is_one();
62
63 let mut scalars_lo = Vec::new();
64 let mut scalars_hi = Vec::new();
65 for (i, scalar) in scalars?.into_iter().enumerate() {
66 if i % 2 == 0 {
67 scalars_lo.push(scalar);
68 } else {
69 scalars_hi.push(scalar);
70 }
71 }
72 let (res_x, res_y) = backend.multi_scalar_mul(&points, &scalars_lo, &scalars_hi, predicate)?;
74 Ok((res_x, res_y))
75}
76
77pub(super) fn embedded_curve_add<F: AcirField>(
78 backend: &impl BlackBoxFunctionSolver<F>,
79 initial_witness: &mut WitnessMap<F>,
80 input1: [FunctionInput<F>; 2],
81 input2: [FunctionInput<F>; 2],
82 predicate: FunctionInput<F>,
83 outputs: (Witness, Witness),
84) -> Result<(), OpcodeResolutionError<F>> {
85 let (res_x, res_y) =
86 execute_embedded_curve_add(backend, initial_witness, input1, input2, predicate)?;
87
88 insert_value(&outputs.0, res_x, initial_witness)?;
89 insert_value(&outputs.1, res_y, initial_witness)?;
90 Ok(())
91}
92
93pub(crate) fn execute_embedded_curve_add<F: AcirField>(
94 backend: &impl BlackBoxFunctionSolver<F>,
95 initial_witness: &WitnessMap<F>,
96 input1: [FunctionInput<F>; 2],
97 input2: [FunctionInput<F>; 2],
98 predicate: FunctionInput<F>,
99) -> Result<(F, F), OpcodeResolutionError<F>> {
100 check_all_or_nothing_pair(BlackBoxFunc::EmbeddedCurveAdd, "Coordinates", input1[0], input1[1])?;
101 check_all_or_nothing_pair(BlackBoxFunc::EmbeddedCurveAdd, "Coordinates", input2[0], input2[1])?;
102
103 let input1_x = input_to_value(initial_witness, input1[0])?;
104 let input1_y = input_to_value(initial_witness, input1[1])?;
105 let input2_x = input_to_value(initial_witness, input2[0])?;
106 let input2_y = input_to_value(initial_witness, input2[1])?;
107 let predicate = input_to_value(initial_witness, predicate)?.is_one();
108 let (res_x, res_y) = backend.ec_add(&input1_x, &input1_y, &input2_x, &input2_y, predicate)?;
109
110 Ok((res_x, res_y))
111}
112
113fn check_all_or_nothing_pair<F: AcirField>(
117 func: BlackBoxFunc,
118 kind: &str,
119 first: FunctionInput<F>,
120 second: FunctionInput<F>,
121) -> Result<(), OpcodeResolutionError<F>> {
122 match (first, second) {
123 (FunctionInput::Witness(_), FunctionInput::Witness(_))
124 | (FunctionInput::Constant(_), FunctionInput::Constant(_)) => Ok(()),
125 _ => Err(OpcodeResolutionError::BlackBoxFunctionFailed(
126 func,
127 format!(
128 "{kind} must be either both witnesses or both constants. Found: {first:?}, {second:?}"
129 ),
130 )),
131 }
132}
133
134#[cfg(test)]
135mod tests {
136 use std::collections::BTreeMap;
137
138 use acir::{
139 AcirField, BlackBoxFunc, FieldElement,
140 circuit::opcodes::FunctionInput,
141 native_types::{Witness, WitnessMap},
142 };
143 use bn254_blackbox_solver::Bn254BlackBoxSolver;
144
145 use super::{execute_embedded_curve_add, execute_multi_scalar_mul};
146 use crate::pwg::OpcodeResolutionError;
147
148 fn generator_y() -> FieldElement {
150 FieldElement::try_from_str("17631683881184975370165255887551781615748388533673675138860")
151 .unwrap()
152 }
153
154 fn witness_map() -> WitnessMap<FieldElement> {
158 WitnessMap::from(BTreeMap::from_iter([
159 (Witness(1), FieldElement::one()),
160 (Witness(2), generator_y()),
161 (Witness(3), FieldElement::one()),
162 (Witness(4), FieldElement::zero()),
163 ]))
164 }
165
166 fn msm(
167 points: [FunctionInput<FieldElement>; 2],
168 scalars: [FunctionInput<FieldElement>; 2],
169 ) -> Result<(FieldElement, FieldElement), OpcodeResolutionError<FieldElement>> {
170 execute_multi_scalar_mul(
171 &Bn254BlackBoxSolver,
172 &witness_map(),
173 &points,
174 &scalars,
175 FunctionInput::Constant(FieldElement::one()),
176 )
177 }
178
179 fn assert_mixed_pair_rejected<T: std::fmt::Debug>(
180 result: Result<T, OpcodeResolutionError<FieldElement>>,
181 expected_func: BlackBoxFunc,
182 ) {
183 match result {
184 Err(OpcodeResolutionError::BlackBoxFunctionFailed(func, message)) => {
185 assert_eq!(func, expected_func);
186 assert!(
187 message.contains("both witnesses or both constants"),
188 "unexpected failure message: {message}"
189 );
190 }
191 other => panic!("expected a mixed constant/witness pair to be rejected, got {other:?}"),
192 }
193 }
194
195 #[test]
196 fn multi_scalar_mul_accepts_uniform_scalar_limbs() {
197 let points = [FunctionInput::Witness(Witness(1)), FunctionInput::Witness(Witness(2))];
198
199 let all_witnesses =
200 msm(points, [FunctionInput::Witness(Witness(3)), FunctionInput::Witness(Witness(4))])
201 .expect("all-witness scalar limbs are valid");
202 let all_constants = msm(
203 points,
204 [
205 FunctionInput::Constant(FieldElement::one()),
206 FunctionInput::Constant(FieldElement::zero()),
207 ],
208 )
209 .expect("all-constant scalar limbs are valid");
210
211 assert_eq!(all_witnesses, all_constants);
213 assert_eq!(all_witnesses, (FieldElement::one(), generator_y()));
214 }
215
216 #[test]
217 fn multi_scalar_mul_rejects_mixed_scalar_limbs() {
218 let points = [FunctionInput::Witness(Witness(1)), FunctionInput::Witness(Witness(2))];
219
220 assert_mixed_pair_rejected(
221 msm(
222 points,
223 [FunctionInput::Constant(FieldElement::one()), FunctionInput::Witness(Witness(4))],
224 ),
225 BlackBoxFunc::MultiScalarMul,
226 );
227 assert_mixed_pair_rejected(
228 msm(
229 points,
230 [FunctionInput::Witness(Witness(3)), FunctionInput::Constant(FieldElement::zero())],
231 ),
232 BlackBoxFunc::MultiScalarMul,
233 );
234 }
235
236 #[test]
237 fn multi_scalar_mul_rejects_mixed_point_coordinates() {
238 let scalars = [FunctionInput::Witness(Witness(3)), FunctionInput::Witness(Witness(4))];
239
240 assert_mixed_pair_rejected(
241 msm(
242 [FunctionInput::Constant(FieldElement::one()), FunctionInput::Witness(Witness(2))],
243 scalars,
244 ),
245 BlackBoxFunc::MultiScalarMul,
246 );
247 }
248
249 #[test]
250 fn embedded_curve_add_rejects_mixed_point_coordinates() {
251 let uniform = [FunctionInput::Witness(Witness(1)), FunctionInput::Witness(Witness(2))];
252 let mixed =
253 [FunctionInput::Constant(FieldElement::one()), FunctionInput::Witness(Witness(2))];
254
255 for (input1, input2) in [(mixed, uniform), (uniform, mixed)] {
256 assert_mixed_pair_rejected(
257 execute_embedded_curve_add(
258 &Bn254BlackBoxSolver,
259 &witness_map(),
260 input1,
261 input2,
262 FunctionInput::Constant(FieldElement::one()),
263 ),
264 BlackBoxFunc::EmbeddedCurveAdd,
265 );
266 }
267 }
268}