use crate::allocator::{Allocator, Atom, NodePtr};
use crate::chia_dialect::ClvmFlags;
use crate::cost::{Cost, check_cost};
use crate::error::EvalErr;
use crate::op_utils::{
MALLOC_COST_PER_BYTE, atom, first, get_args, get_varargs, int_atom, mod_group_order,
new_atom_and_cost, nilp, rest,
};
use crate::reduction::{Reduction, Response};
use chia_bls::{
G1Element, G2Element, PublicKey, aggregate_pairing, aggregate_verify, hash_to_g1_with_dst,
hash_to_g2_with_dst,
};
const BLS_G1_SUBTRACT_BASE_COST: Cost = 101094;
const BLS_G1_SUBTRACT_COST_PER_ARG: Cost = 1343980;
const BLS_G1_MULTIPLY_BASE_COST: Cost = 705500;
const BLS_G1_MULTIPLY_COST_PER_BYTE: Cost = 10;
const BLS_G1_NEGATE_BASE_COST: Cost = 1396 - 480;
const BLS_G2_ADD_BASE_COST: Cost = 80000;
const BLS_G2_ADD_COST_PER_ARG: Cost = 1950000;
const BLS_G2_SUBTRACT_BASE_COST: Cost = 80000;
const BLS_G2_SUBTRACT_COST_PER_ARG: Cost = 1950000;
const BLS_G2_MULTIPLY_BASE_COST: Cost = 2100000;
const BLS_G2_MULTIPLY_COST_PER_BYTE: Cost = 5;
const BLS_G2_NEGATE_BASE_COST: Cost = 2164 - 960;
const BLS_MAP_TO_G1_BASE_COST: Cost = 195000;
const BLS_MAP_TO_G1_COST_PER_BYTE: Cost = 4;
const BLS_MAP_TO_G1_COST_PER_DST_BYTE: Cost = 4;
const BLS_MAP_TO_G2_BASE_COST: Cost = 815000;
const BLS_MAP_TO_G2_COST_PER_BYTE: Cost = 4;
const BLS_MAP_TO_G2_COST_PER_DST_BYTE: Cost = 4;
const BLS_PAIRING_BASE_COST: Cost = 3000000;
const BLS_PAIRING_COST_PER_ARG: Cost = 1200000;
const DST_G2: &[u8; 43] = b"BLS_SIG_BLS12381G2_XMD:SHA-256_SSWU_RO_AUG_";
pub fn op_bls_g1_subtract(
a: &mut Allocator,
mut input: NodePtr,
max_cost: Cost,
_flags: ClvmFlags,
) -> Response {
let mut cost = BLS_G1_SUBTRACT_BASE_COST;
check_cost(cost, max_cost)?;
let mut total = G1Element::default();
let mut is_first = true;
while let Some((arg, rest)) = a.next(input) {
input = rest;
let point = a.g1(arg)?;
cost += BLS_G1_SUBTRACT_COST_PER_ARG;
check_cost(cost, max_cost)?;
if is_first {
total = point;
} else {
total -= &point;
};
is_first = false;
}
Ok(Reduction(
cost + 48 * MALLOC_COST_PER_BYTE,
a.new_g1(total)?,
))
}
pub fn op_bls_g1_multiply(
a: &mut Allocator,
input: NodePtr,
max_cost: Cost,
_flags: ClvmFlags,
) -> Response {
let [point, scalar] = get_args::<2>(a, input, "g1_multiply")?;
let mut cost = BLS_G1_MULTIPLY_BASE_COST;
check_cost(cost, max_cost)?;
let mut total = a.g1(point)?;
let (scalar, scalar_len) = int_atom(a, scalar, "g1_multiply")?;
if scalar_len > 1024 {
return Err(EvalErr::InvalidOpArg(input, "g1_multiply".to_string()));
}
cost += scalar_len as Cost * BLS_G1_MULTIPLY_COST_PER_BYTE;
check_cost(cost, max_cost)?;
let scalar = mod_group_order(scalar);
total.scalar_multiply(scalar.to_bytes_be().1.as_slice());
Ok(Reduction(
cost + 48 * MALLOC_COST_PER_BYTE,
a.new_g1(total)?,
))
}
pub fn op_bls_g1_negate(
a: &mut Allocator,
input: NodePtr,
_max_cost: Cost,
flags: ClvmFlags,
) -> Response {
let strict = !flags.contains(ClvmFlags::RELAXED_BLS);
let [point] = get_args::<1>(a, input, "g1_negate")?;
let mut blob: [u8; 48] = atom(a, point, "G1 atom").and_then(|blob| {
blob.as_ref().try_into().map_err(|_| {
EvalErr::InvalidOpArg(point, "atom is not a G1 size, 48 bytes".to_string())
})
})?;
if strict {
a.validate_g1(point, blob)?;
}
if (blob[0] & 0xe0) == 0xc0 {
Ok(Reduction(
BLS_G1_NEGATE_BASE_COST + 48 * MALLOC_COST_PER_BYTE,
point,
))
} else {
blob[0] ^= 0x20;
if strict {
a.add_validated_g1(blob);
}
new_atom_and_cost(a, BLS_G1_NEGATE_BASE_COST, &blob)
}
}
pub fn op_bls_g2_add(
a: &mut Allocator,
mut input: NodePtr,
max_cost: Cost,
_flags: ClvmFlags,
) -> Response {
let mut cost = BLS_G2_ADD_BASE_COST;
check_cost(cost, max_cost)?;
let mut total = G2Element::default();
while let Some((arg, rest)) = a.next(input) {
input = rest;
let point = a.g2(arg)?;
cost += BLS_G2_ADD_COST_PER_ARG;
check_cost(cost, max_cost)?;
total += &point;
}
Ok(Reduction(
cost + 96 * MALLOC_COST_PER_BYTE,
a.new_g2(total)?,
))
}
pub fn op_bls_g2_subtract(
a: &mut Allocator,
mut input: NodePtr,
max_cost: Cost,
_flags: ClvmFlags,
) -> Response {
let mut cost = BLS_G2_SUBTRACT_BASE_COST;
check_cost(cost, max_cost)?;
let mut total = G2Element::default();
let mut is_first = true;
while let Some((arg, rest)) = a.next(input) {
input = rest;
let point = a.g2(arg)?;
cost += BLS_G2_SUBTRACT_COST_PER_ARG;
check_cost(cost, max_cost)?;
if is_first {
total = point;
} else {
total -= &point;
};
is_first = false;
}
Ok(Reduction(
cost + 96 * MALLOC_COST_PER_BYTE,
a.new_g2(total)?,
))
}
pub fn op_bls_g2_multiply(
a: &mut Allocator,
input: NodePtr,
max_cost: Cost,
_flags: ClvmFlags,
) -> Response {
let [point, scalar] = get_args::<2>(a, input, "g2_multiply")?;
let mut cost = BLS_G2_MULTIPLY_BASE_COST;
check_cost(cost, max_cost)?;
let mut total = a.g2(point)?;
let (scalar, scalar_len) = int_atom(a, scalar, "g2_multiply")?;
if scalar_len > 1024 {
return Err(EvalErr::InvalidOpArg(input, "g2_multiply".to_string()));
}
cost += scalar_len as Cost * BLS_G2_MULTIPLY_COST_PER_BYTE;
check_cost(cost, max_cost)?;
let scalar = mod_group_order(scalar);
total.scalar_multiply(scalar.to_bytes_be().1.as_slice());
Ok(Reduction(
cost + 96 * MALLOC_COST_PER_BYTE,
a.new_g2(total)?,
))
}
pub fn op_bls_g2_negate(
a: &mut Allocator,
input: NodePtr,
_max_cost: Cost,
flags: ClvmFlags,
) -> Response {
let strict = !flags.contains(ClvmFlags::RELAXED_BLS);
let [point] = get_args::<1>(a, input, "g2_negate")?;
let mut blob: [u8; 96] = atom(a, point, "G2 atom").and_then(|blob| {
blob.as_ref()
.try_into()
.map_err(|_| EvalErr::InvalidOpArg(point, "atom is not G2 size, 96 bytes".to_string()))
})?;
if strict {
a.validate_g2(point, blob)?;
}
if (blob[0] & 0xe0) == 0xc0 {
Ok(Reduction(
BLS_G2_NEGATE_BASE_COST + 96 * MALLOC_COST_PER_BYTE,
point,
))
} else {
blob[0] ^= 0x20;
if strict {
a.add_validated_g2(blob);
}
new_atom_and_cost(a, BLS_G2_NEGATE_BASE_COST, &blob)
}
}
pub fn op_bls_map_to_g1(
a: &mut Allocator,
input: NodePtr,
max_cost: Cost,
_flags: ClvmFlags,
) -> Response {
let ([msg, dst], argc) = get_varargs::<2>(a, input, "g1_map")?;
if !(1..=2).contains(&argc) {
Err(EvalErr::InvalidOpArg(
input,
format!("g1_map takes exactly 1 or 2 arguments, got {argc}"),
))?;
}
let mut cost: Cost = BLS_MAP_TO_G1_BASE_COST;
check_cost(cost, max_cost)?;
let msg = atom(a, msg, "g1_map")?;
cost += msg.as_ref().len() as Cost * BLS_MAP_TO_G1_COST_PER_BYTE;
check_cost(cost, max_cost)?;
let dst = if argc == 2 {
atom(a, dst, "g1_map")?
} else {
Atom::Borrowed(b"BLS_SIG_BLS12381G1_XMD:SHA-256_SSWU_RO_AUG_".as_slice())
};
cost += dst.as_ref().len() as Cost * BLS_MAP_TO_G1_COST_PER_DST_BYTE;
check_cost(cost, max_cost)?;
let point = hash_to_g1_with_dst(msg.as_ref(), dst.as_ref());
Ok(Reduction(
cost + 48 * MALLOC_COST_PER_BYTE,
a.new_g1(point)?,
))
}
pub fn op_bls_map_to_g2(
a: &mut Allocator,
input: NodePtr,
max_cost: Cost,
_flags: ClvmFlags,
) -> Response {
let ([msg, dst], argc) = get_varargs::<2>(a, input, "g2_map")?;
if !(1..=2).contains(&argc) {
Err(EvalErr::InvalidOpArg(
input,
format!("g2_map takes exactly 1 or 2 arguments, got {argc}"),
))?;
}
let mut cost: Cost = BLS_MAP_TO_G2_BASE_COST;
check_cost(cost, max_cost)?;
let msg = atom(a, msg, "g2_map")?;
cost += msg.as_ref().len() as Cost * BLS_MAP_TO_G2_COST_PER_BYTE;
let dst = if argc == 2 {
atom(a, dst, "g2_map")?
} else {
Atom::Borrowed(DST_G2.as_slice())
};
cost += dst.as_ref().len() as Cost * BLS_MAP_TO_G2_COST_PER_DST_BYTE;
check_cost(cost, max_cost)?;
let point = hash_to_g2_with_dst(msg.as_ref(), dst.as_ref());
Ok(Reduction(
cost + 96 * MALLOC_COST_PER_BYTE,
a.new_g2(point)?,
))
}
pub fn op_bls_pairing_identity(
a: &mut Allocator,
input: NodePtr,
max_cost: Cost,
_flags: ClvmFlags,
) -> Response {
let mut cost = BLS_PAIRING_BASE_COST;
check_cost(cost, max_cost)?;
let mut items = Vec::<(G1Element, G2Element)>::new();
let mut args = input;
while !nilp(a, args) {
cost += BLS_PAIRING_COST_PER_ARG;
check_cost(cost, max_cost)?;
let g1 = a.g1(first(a, args)?)?;
args = rest(a, args)?;
let g2 = a.g2(first(a, args)?)?;
args = rest(a, args)?;
items.push((g1, g2));
}
if !aggregate_pairing(items) {
Err(EvalErr::BLSPairingIdentityFailed(input))?
} else {
Ok(Reduction(cost, a.nil()))
}
}
pub fn op_bls_verify(
a: &mut Allocator,
input: NodePtr,
max_cost: Cost,
_flags: ClvmFlags,
) -> Response {
let mut cost = BLS_PAIRING_BASE_COST;
check_cost(cost, max_cost)?;
let mut args = input;
let signature = a.g2(first(a, args)?)?;
args = rest(a, args)?;
let mut items = Vec::<(PublicKey, Atom)>::new();
while !nilp(a, args) {
let pk = a.g1(first(a, args)?)?;
args = rest(a, args)?;
let msg = atom(a, first(a, args)?, "bls_verify message")?;
args = rest(a, args)?;
cost += BLS_PAIRING_COST_PER_ARG;
cost += msg.as_ref().len() as Cost * BLS_MAP_TO_G2_COST_PER_BYTE;
cost += DST_G2.len() as Cost * BLS_MAP_TO_G2_COST_PER_DST_BYTE;
check_cost(cost, max_cost)?;
items.push((pk, msg));
}
if !aggregate_verify(&signature, items) {
Err(EvalErr::BLSVerifyFailed(input))?
} else {
Ok(Reduction(cost, a.nil()))
}
}