use alloc::vec::Vec;
use core::cell::Cell;
use bytecheck::CheckBytes;
use c_kzg::{Bytes32 as KzgBytes32, Bytes48};
use dusk_bytes::DeserializableSlice;
use dusk_core::BlsScalar;
use dusk_core::groth16::bn254::{Bn254, G1Projective};
use dusk_core::groth16::serialize::CanonicalDeserialize;
use dusk_core::groth16::{
Groth16, PreparedVerifyingKey, Proof as Groth16Proof,
};
use dusk_core::plonk::{PlonkVersion, Proof as PlonkProof, Verifier};
use dusk_core::signatures::bls::{
self as bls, BlsVersion, MultisigSignature, PublicKey as BlsPublicKey,
Signature as BlsSignature,
};
use dusk_core::signatures::schnorr::{
PublicKey as SchnorrPublicKey, Signature as SchnorrSignature,
};
use dusk_core::transfer::data::BlobData;
use dusk_poseidon::{Domain, Hash as PoseidonHash};
use rkyv::ser::serializers::AllocSerializer;
use rkyv::validation::validators::DefaultValidator;
use rkyv::{Archive, Deserialize, Serialize};
use secp256k1::{Message, Secp256k1, ecdsa::RecoverableSignature};
use sha2::{Digest as Sha2Digest, Sha256};
use sha3::Keccak256;
use tracing::warn;
use crate::cache;
thread_local! {
static PLONK_VERSION: Cell<PlonkVersion> = const { Cell::new(PlonkVersion::V2) };
static HARD_FORK: Cell<HardFork> = const { Cell::new(HardFork::PreFork) };
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub enum HardFork {
PreFork,
Aegis,
}
impl HardFork {
pub fn bls_version(&self) -> BlsVersion {
match self {
HardFork::Aegis => BlsVersion::V2,
HardFork::PreFork => BlsVersion::V1,
}
}
}
#[derive(Debug)]
pub struct PlonkVersionGuard {
prev: PlonkVersion,
}
impl Drop for PlonkVersionGuard {
fn drop(&mut self) {
PLONK_VERSION.with(|m| m.set(self.prev));
}
}
pub fn plonk_version() -> PlonkVersion {
PLONK_VERSION.with(|m| m.get())
}
pub fn set_plonk_version(version: PlonkVersion) -> PlonkVersionGuard {
let prev = PLONK_VERSION.with(|m| {
let prev = m.get();
m.set(version);
prev
});
PlonkVersionGuard { prev }
}
#[derive(Debug)]
pub struct HardForkGuard {
prev: HardFork,
}
impl Drop for HardForkGuard {
fn drop(&mut self) {
HARD_FORK.with(|m| m.set(self.prev));
}
}
pub fn hard_fork() -> HardFork {
HARD_FORK.with(|m| m.get())
}
pub fn set_hard_fork(hard_fork: HardFork) -> HardForkGuard {
let prev = HARD_FORK.with(|m| {
let prev = m.get();
m.set(hard_fork);
prev
});
HardForkGuard { prev }
}
pub fn hash(bytes: Vec<u8>) -> BlsScalar {
BlsScalar::hash_to_scalar(&bytes[..])
}
pub fn poseidon_hash(scalars: Vec<BlsScalar>) -> BlsScalar {
PoseidonHash::digest(Domain::Other, &scalars)[0]
}
pub fn verify_plonk_with_version(
version: PlonkVersion,
verifier_data: Vec<u8>,
proof: Vec<u8>,
public_inputs: Vec<BlsScalar>,
) -> bool {
let verifier = match Verifier::try_from_bytes(verifier_data) {
Ok(v) => v,
Err(e) => {
warn!("vm: couldn't deserialize plonk verifier: {e:?}");
return false;
}
};
let proof = match PlonkProof::from_slice(&proof) {
Ok(p) => p,
Err(e) => {
warn!("vm: couldn't deserialize plonk proof: {e:?}");
return false;
}
};
let result =
verifier.verify_with_version(&proof, &public_inputs[..], version);
match result {
Ok(_) => true,
Err(e) => {
warn!("vm: plonk verification failed ({version:?}): {e:?}");
false
}
}
}
fn plonk_cache_key(
version: PlonkVersion,
arg_buf: &[u8],
) -> [u8; blake2b_simd::OUTBYTES] {
let mut state = blake2b_simd::Params::new()
.hash_length(blake2b_simd::OUTBYTES)
.to_state();
let cache_tag = match version {
PlonkVersion::V1 => 0,
PlonkVersion::V2 => 1,
PlonkVersion::V3 => 2,
_ => u8::MAX,
};
state.update(&[cache_tag]);
state.update(arg_buf);
*state.finalize().as_array()
}
fn bls_cache_key(
hard_fork: HardFork,
arg_buf: &[u8],
) -> [u8; blake2b_simd::OUTBYTES] {
let mut state = blake2b_simd::Params::new()
.hash_length(blake2b_simd::OUTBYTES)
.to_state();
let cache_tag = match hard_fork {
HardFork::PreFork => 0u8,
HardFork::Aegis => 1u8,
};
state.update(&[cache_tag]);
state.update(arg_buf);
*state.finalize().as_array()
}
pub fn verify_groth16_bn254(
pvk: Vec<u8>,
proof: Vec<u8>,
inputs: Vec<u8>,
) -> bool {
let pvk = match PreparedVerifyingKey::deserialize_uncompressed(&pvk[..]) {
Ok(v) => v,
Err(e) => {
warn!("vm: couldn't deserialize groth16 verifiying key: {e}");
return false;
}
};
let proof = match Groth16Proof::deserialize_compressed(&proof[..]) {
Ok(p) => p,
Err(e) => {
warn!("vm: couldn't deserialize groth16 proof: {e}");
return false;
}
};
let inputs = match G1Projective::deserialize_compressed(&inputs[..]) {
Ok(i) => i,
Err(e) => {
warn!("vm: couldn't deserialize groth16 inputs: {e}");
return false;
}
};
match Groth16::<Bn254>::verify_proof_with_prepared_inputs(
&pvk, &proof, &inputs,
) {
Ok(valid) => valid,
Err(e) => {
warn!("vm: couldn't verify groth16: {e}");
false
}
}
}
pub fn verify_schnorr(
msg: BlsScalar,
pk: SchnorrPublicKey,
sig: SchnorrSignature,
) -> bool {
pk.verify(&sig, msg).is_ok()
}
pub fn verify_bls(msg: Vec<u8>, pk: BlsPublicKey, sig: BlsSignature) -> bool {
bls::verify(&pk, &sig, &msg, hard_fork().bls_version()).is_ok()
}
pub fn verify_bls_multisig(
msg: Vec<u8>,
keys: Vec<BlsPublicKey>,
sig: MultisigSignature,
) -> bool {
if keys.is_empty() {
warn!("vm: bls multisig verification requires at least one key");
return false;
}
let bls_version = hard_fork().bls_version();
let akey = match bls::aggregate(&keys, bls_version) {
Ok(k) => k,
Err(e) => {
warn!("vm: couldn't aggregate bls public-keys due to {e}");
return false;
}
};
bls::verify_multisig(&akey, &sig, &msg, bls_version).is_ok()
}
pub fn keccak256(bytes: Vec<u8>) -> [u8; 32] {
let mut hasher = Keccak256::new();
hasher.update(bytes.as_slice());
hasher.finalize().into()
}
pub fn sha256(bytes: Vec<u8>) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(bytes.as_slice());
hasher.finalize().into()
}
pub fn verify_kzg_proof(
commitment: [u8; 48],
z: [u8; 32],
y: [u8; 32],
proof: [u8; 48],
) -> bool {
let settings = BlobData::eth_kzg_settings(None);
let commitment = Bytes48::new(commitment);
let z = KzgBytes32::new(z);
let y = KzgBytes32::new(y);
let proof = Bytes48::new(proof);
match settings.verify_kzg_proof(&commitment, &z, &y, &proof) {
Ok(valid) => valid,
Err(e) => {
warn!("vm: kzg proof verification failed: {e}");
false
}
}
}
pub fn secp256k1_recover(
msg_hash: [u8; 32],
sig: [u8; 65],
) -> Option<[u8; 65]> {
let v_raw = sig[64];
let v = match v_raw {
0 | 1 => v_raw as i32,
27 | 28 => (v_raw - 27) as i32,
_ => {
warn!("vm: secp256k1 recovery: invalid v byte {v_raw}");
return None;
}
};
let rec_id = match secp256k1::ecdsa::RecoveryId::try_from(v) {
Ok(id) => id,
Err(e) => {
warn!("vm: secp256k1 recovery: invalid recovery id {v} ({e})");
return None;
}
};
let sig = match RecoverableSignature::from_compact(&sig[0..64], rec_id) {
Ok(sig) => sig,
Err(e) => {
warn!("vm: secp256k1 recovery: invalid signature ({e})");
return None;
}
};
let msg = Message::from_digest(msg_hash);
let secp = Secp256k1::new();
let pk = match secp.recover_ecdsa(msg, &sig) {
Ok(pk) => pk,
Err(e) => {
warn!("vm: secp256k1 recovery failed ({e})");
return None;
}
};
Some(pk.serialize_uncompressed())
}
fn write_to_arg_buf<R>(arg_buf: &mut [u8], result: &R) -> u32
where
R: Serialize<AllocSerializer<1024>>,
{
let bytes = rkyv::to_bytes::<_, 1024>(result).unwrap();
arg_buf[..bytes.len()].copy_from_slice(&bytes);
bytes.len() as u32
}
fn wrap_host_query<A, R, F>(
arg_buf: &mut [u8],
arg_len: u32,
name: &str,
fallback: &R,
closure: F,
) -> u32
where
F: FnOnce(A) -> R,
A: Archive,
A::Archived: for<'a> CheckBytes<DefaultValidator<'a>>
+ Deserialize<A, rkyv::Infallible>,
R: Serialize<AllocSerializer<1024>>,
{
let Some(root) =
rkyv::check_archived_root::<A>(&arg_buf[..arg_len as usize]).ok()
else {
warn!("vm: invalid archived data in {name}");
return write_to_arg_buf(arg_buf, fallback);
};
let arg: A = root.deserialize(&mut rkyv::Infallible).unwrap();
let result = closure(arg);
write_to_arg_buf(arg_buf, &result)
}
pub(crate) fn host_hash(arg_buf: &mut [u8], arg_len: u32) -> u32 {
wrap_host_query(arg_buf, arg_len, "host_hash", &BlsScalar::default(), hash)
}
pub(crate) fn host_poseidon_hash(arg_buf: &mut [u8], arg_len: u32) -> u32 {
wrap_host_query(
arg_buf,
arg_len,
"host_poseidon_hash",
&BlsScalar::default(),
poseidon_hash,
)
}
pub(crate) fn host_verify_plonk(arg_buf: &mut [u8], arg_len: u32) -> u32 {
let version = plonk_version();
let hash = plonk_cache_key(version, &arg_buf[..arg_len as usize]);
let cached = cache::get_plonk_verification(hash);
wrap_host_query(
arg_buf,
arg_len,
"host_verify_plonk",
&false,
|(vd, proof, pis)| {
let is_valid = cached.unwrap_or_else(|| {
verify_plonk_with_version(version, vd, proof, pis)
});
cache::put_plonk_verification(hash, is_valid);
is_valid
},
)
}
pub(crate) fn host_verify_groth16_bn254(
arg_buf: &mut [u8],
arg_len: u32,
) -> u32 {
let hash = *blake2b_simd::blake2b(&arg_buf[..arg_len as usize]).as_array();
let cached = cache::get_groth16_verification(hash);
wrap_host_query(
arg_buf,
arg_len,
"host_verify_groth16_bn254",
&false,
|(pvk, proof, inputs)| {
let is_valid = cached
.unwrap_or_else(|| verify_groth16_bn254(pvk, proof, inputs));
cache::put_groth16_verification(hash, is_valid);
is_valid
},
)
}
pub(crate) fn host_verify_schnorr(arg_buf: &mut [u8], arg_len: u32) -> u32 {
wrap_host_query(
arg_buf,
arg_len,
"host_verify_schnorr",
&false,
|(msg, pk, sig)| verify_schnorr(msg, pk, sig),
)
}
pub(crate) fn host_verify_bls(arg_buf: &mut [u8], arg_len: u32) -> u32 {
let current_hard_fork = hard_fork();
let hash = bls_cache_key(current_hard_fork, &arg_buf[..arg_len as usize]);
let cached = cache::get_bls_verification(hash);
wrap_host_query(
arg_buf,
arg_len,
"host_verify_bls",
&false,
|(msg, pk, sig)| {
let is_valid = cached.unwrap_or_else(|| verify_bls(msg, pk, sig));
cache::put_bls_verification(hash, is_valid);
is_valid
},
)
}
pub(crate) fn host_verify_bls_multisig(
arg_buf: &mut [u8],
arg_len: u32,
) -> u32 {
wrap_host_query(
arg_buf,
arg_len,
"host_verify_bls_multisig",
&false,
|(msg, keys, sig)| verify_bls_multisig(msg, keys, sig),
)
}
pub(crate) fn host_keccak256(arg_buf: &mut [u8], arg_len: u32) -> u32 {
wrap_host_query(arg_buf, arg_len, "host_keccak256", &[0u8; 32], keccak256)
}
pub(crate) fn host_sha256(arg_buf: &mut [u8], arg_len: u32) -> u32 {
wrap_host_query(arg_buf, arg_len, "host_sha256", &[0u8; 32], sha256)
}
pub(crate) fn host_verify_kzg_proof(arg_buf: &mut [u8], arg_len: u32) -> u32 {
wrap_host_query(
arg_buf,
arg_len,
"host_verify_kzg_proof",
&false,
|(commitment, z, y, proof)| verify_kzg_proof(commitment, z, y, proof),
)
}
pub(crate) fn host_secp256k1_recover(arg_buf: &mut [u8], arg_len: u32) -> u32 {
wrap_host_query(
arg_buf,
arg_len,
"host_secp256k1_recover",
&Option::<[u8; 65]>::None,
|(msg_hash, sig)| secp256k1_recover(msg_hash, sig),
)
}