1use acir::{
2 AcirField,
3 circuit::opcodes::{BlackBoxFuncCall, FunctionInput},
4 native_types::{Witness, WitnessMap},
5};
6use acvm_blackbox_solver::{blake2s, blake3, keccakf1600};
7use itertools::Itertools;
8
9use self::{aes128::solve_aes128_encryption_opcode, hash::solve_poseidon2_permutation_opcode};
10
11use super::{OpcodeNotSolvable, OpcodeResolutionError, insert_value};
12use crate::{
13 BlackBoxFunctionSolver,
14 pwg::{check_bit_size, input_to_value},
15};
16
17pub(crate) mod aes128;
18pub(crate) mod embedded_curve_ops;
19pub(crate) mod hash;
20mod logic;
21mod range;
22pub(crate) mod signature;
23pub(crate) mod utils;
24
25use embedded_curve_ops::{embedded_curve_add, multi_scalar_mul};
26use hash::{solve_generic_256_hash_opcode, solve_sha_256_permutation_opcode};
28use logic::{and, xor};
29pub(crate) use range::solve_range_opcode;
30use signature::ecdsa::{secp256k1_prehashed, secp256r1_prehashed};
31
32fn first_missing_assignment<F>(
36 witness_assignments: &WitnessMap<F>,
37 inputs: &[FunctionInput<F>],
38) -> Option<Witness> {
39 inputs.iter().find_map(|input| {
40 if let FunctionInput::Witness(witness) = input {
41 if witness_assignments.contains_key(witness) { None } else { Some(*witness) }
42 } else {
43 None
44 }
45 })
46}
47
48fn contains_all_inputs<F>(
50 witness_assignments: &WitnessMap<F>,
51 inputs: &[FunctionInput<F>],
52) -> bool {
53 first_missing_assignment(witness_assignments, inputs).is_none()
54}
55
56pub(crate) fn solve<F: AcirField>(
70 backend: &impl BlackBoxFunctionSolver<F>,
71 initial_witness: &mut WitnessMap<F>,
72 bb_func: &BlackBoxFuncCall<F>,
73) -> Result<(), OpcodeResolutionError<F>> {
74 let inputs = bb_func.get_inputs_vec();
75 if !contains_all_inputs(initial_witness, &inputs) {
76 let unassigned_witness = first_missing_assignment(initial_witness, &inputs)
77 .expect("Some assignments must be missing because it does not contains all inputs");
78 return Err(OpcodeResolutionError::OpcodeNotSolvable(
79 OpcodeNotSolvable::MissingAssignment(unassigned_witness.witness_index()),
80 ));
81 }
82
83 match bb_func {
84 BlackBoxFuncCall::AES128Encrypt { inputs, iv, key, outputs } => {
85 solve_aes128_encryption_opcode(initial_witness, inputs, iv, key, outputs)
86 }
87 BlackBoxFuncCall::AND { lhs, rhs, num_bits, output } => {
88 and(initial_witness, lhs, rhs, *num_bits, output)
89 }
90 BlackBoxFuncCall::XOR { lhs, rhs, num_bits, output } => {
91 xor(initial_witness, lhs, rhs, *num_bits, output)
92 }
93 BlackBoxFuncCall::RANGE { input, num_bits } => {
94 solve_range_opcode(initial_witness, input, *num_bits)
95 }
96 BlackBoxFuncCall::Blake2s { outputs, .. } => {
97 let inputs = bb_func.get_inputs_vec();
98 solve_generic_256_hash_opcode(initial_witness, &inputs, None, outputs, blake2s)
99 }
100 BlackBoxFuncCall::Blake3 { outputs, .. } => {
101 let inputs = bb_func.get_inputs_vec();
102 solve_generic_256_hash_opcode(initial_witness, &inputs, None, outputs, blake3)
103 }
104 BlackBoxFuncCall::Keccakf1600 { inputs, outputs } => {
105 let mut state = [0; 25];
106 for (it, input) in state.iter_mut().zip_eq(inputs.as_ref()) {
107 let witness_assignment = input_to_value(initial_witness, *input)?;
108 check_bit_size(witness_assignment, 64)?;
109 *it = witness_assignment
110 .try_to_u64()
111 .expect("value was just checked to fit in 64 bits");
112 }
113 let output_state = keccakf1600(state)?;
114 for (output_witness, value) in outputs.iter().zip_eq(output_state) {
115 insert_value(output_witness, F::from(u128::from(value)), initial_witness)?;
116 }
117 Ok(())
118 }
119 BlackBoxFuncCall::EcdsaSecp256k1 {
120 public_key_x,
121 public_key_y,
122 signature,
123 hashed_message: message,
124 output,
125 predicate,
126 } => secp256k1_prehashed(
127 initial_witness,
128 public_key_x,
129 public_key_y,
130 signature,
131 message.as_ref(),
132 predicate,
133 *output,
134 ),
135 BlackBoxFuncCall::EcdsaSecp256r1 {
136 public_key_x,
137 public_key_y,
138 signature,
139 hashed_message: message,
140 output,
141 predicate,
142 } => secp256r1_prehashed(
143 initial_witness,
144 public_key_x,
145 public_key_y,
146 signature,
147 message.as_ref(),
148 predicate,
149 *output,
150 ),
151 BlackBoxFuncCall::MultiScalarMul { points, scalars, outputs, predicate } => {
152 multi_scalar_mul(backend, initial_witness, points, scalars, *predicate, *outputs)
153 }
154 BlackBoxFuncCall::EmbeddedCurveAdd { input1, input2, outputs, predicate } => {
155 embedded_curve_add(backend, initial_witness, **input1, **input2, *predicate, *outputs)
156 }
157 BlackBoxFuncCall::RecursiveAggregation { .. } => Ok(()),
159 BlackBoxFuncCall::Sha256Compression { inputs, hash_values, outputs } => {
160 solve_sha_256_permutation_opcode(initial_witness, inputs, hash_values, outputs)
161 }
162 BlackBoxFuncCall::Poseidon2Permutation { inputs, outputs } => {
163 solve_poseidon2_permutation_opcode(backend, initial_witness, inputs, outputs)
164 }
165 }
166}