acvm/pwg/blackbox/
embedded_curve_ops.rs

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 the resulting point into the witness map
22    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    // Call the backend's multi-scalar multiplication function
73    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
113/// Checks that the two halves of an input pair are either both witnesses or both constants,
114/// erroring otherwise. `kind` names the pair in the error message, e.g. "Coordinates" for a
115/// point's `(x, y)` or "Scalar limbs" for a scalar's `(lo, hi)`.
116fn 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    /// `y` coordinate of the Grumpkin generator, whose `x` coordinate is 1.
149    fn generator_y() -> FieldElement {
150        FieldElement::try_from_str("17631683881184975370165255887551781615748388533673675138860")
151            .unwrap()
152    }
153
154    /// Witness map holding the generator's `y` coordinate in `Witness(2)` and the scalar
155    /// `1` split into limbs `lo = 1` (`Witness(3)`) and `hi = 0` (`Witness(4)`).
156    /// `Witness(1)` holds the generator's `x` coordinate, which is `1`.
157    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        // `1 * G == G`, whichever way the scalar limbs are declared.
212        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}