mod glv;
mod secp256k1;
mod short_weierstrass;
use alloc::vec::Vec;
use miden_core::{
Felt, ZERO,
deferred::{
DeferredContext, DeferredError, Digest, Node, NodeType, Payload, Precompile,
PrecompileError, TRUE_DIGEST, Tag, precompile_id,
},
};
use self::secp256k1::Secp256k1;
pub use self::{
glv::{SECP256K1_BETA, SECP256K1_LAMBDA, glv_decompose, phi_generator, scalar_mul_mod_n},
secp256k1::{SECP256K1_GENERATOR_X, SECP256K1_GENERATOR_Y, SECP256K1_ID},
};
use crate::math::uint::{Limbs, UintDomain, UintPrecompile, UintSpec};
pub const K1_A_PTR: u32 = 8;
pub const K1_B_PTR: u32 = 9;
pub const K1_BETA_PTR: u32 = 10;
pub const K1_LAMBDA_PTR: u32 = 11;
pub const K1_GROUP_PTR: u32 = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct CurveCoefficient {
pub ptr: u32,
pub bound_ptr: u32,
pub value: Limbs,
}
pub fn curve_coefficients() -> [CurveCoefficient; 2] {
[
CurveCoefficient {
ptr: CurveId::Secp256k1.a_ptr(),
bound_ptr: CurveId::Secp256k1.base_domain().bound_ptr(),
value: <Secp256k1 as ShortWeierstrassSpec>::A,
},
CurveCoefficient {
ptr: CurveId::Secp256k1.b_ptr(),
bound_ptr: CurveId::Secp256k1.base_domain().bound_ptr(),
value: <Secp256k1 as ShortWeierstrassSpec>::B,
},
]
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Endomorphism {
pub beta_ptr: u32,
pub beta: Limbs,
pub lambda_ptr: u32,
pub lambda: Limbs,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CurvePoint {
Identity,
Affine { x: Limbs, y: Limbs },
}
pub trait CurveSpec: Sized + 'static {
const ID: Felt;
type BaseField: UintSpec;
type ScalarField: UintSpec;
const GENERATOR_X: Limbs;
const GENERATOR_Y: Limbs;
fn generator() -> CurvePoint {
Self::point_from_affine(Self::GENERATOR_X, Self::GENERATOR_Y)
.expect("curve generator coordinates must be valid")
}
fn point_from_affine(x: Limbs, y: Limbs) -> Result<CurvePoint, PrecompileError>;
fn canonical_point(point: CurvePoint) -> Result<CurvePoint, PrecompileError> {
match point {
CurvePoint::Identity => Ok(CurvePoint::Identity),
CurvePoint::Affine { x, y } => Self::point_from_affine(x, y),
}
}
fn is_on_curve(point: &CurvePoint) -> bool {
Self::canonical_point(*point).is_ok()
}
fn add(lhs: CurvePoint, rhs: CurvePoint) -> Result<CurvePoint, PrecompileError>;
fn neg(point: CurvePoint) -> Result<CurvePoint, PrecompileError>;
fn sub(lhs: CurvePoint, rhs: CurvePoint) -> Result<CurvePoint, PrecompileError> {
let rhs = Self::neg(rhs)?;
Self::add(lhs, rhs)
}
fn mul_scalar(point: CurvePoint, scalar: Limbs) -> Result<CurvePoint, PrecompileError> {
debug_assert!(Self::is_on_curve(&point));
debug_assert!(Self::ScalarField::is_canonical(&scalar));
let Some(highest_limb) = scalar.iter().rposition(|&limb| limb != 0) else {
return Ok(CurvePoint::Identity);
};
let highest_bit =
highest_limb * 32 + (u32::BITS - 1 - scalar[highest_limb].leading_zeros()) as usize;
let mut acc = CurvePoint::Identity;
let mut base = point;
for bit_index in 0..=highest_bit {
let limb = scalar[bit_index / 32];
if ((limb >> (bit_index % 32)) & 1) == 1 {
acc = Self::add(acc, base)?;
}
if bit_index != highest_bit {
base = Self::add(base, base)?;
}
}
Ok(acc)
}
}
pub trait ShortWeierstrassSpec: CurveSpec {
const A: Limbs;
const B: Limbs;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CurveId {
Secp256k1,
}
impl CurveId {
pub const ALL: [Self; 1] = [Self::Secp256k1];
pub fn from_id(id: Felt) -> Option<Self> {
match id {
id if id == <Secp256k1 as CurveSpec>::ID => Some(Self::Secp256k1),
_ => None,
}
}
pub fn id(self) -> Felt {
match self {
Self::Secp256k1 => SECP256K1_ID,
}
}
pub const fn group_ptr(self) -> u32 {
match self {
Self::Secp256k1 => K1_GROUP_PTR,
}
}
pub const fn from_group_ptr(ptr: u32) -> Option<Self> {
match ptr {
K1_GROUP_PTR => Some(Self::Secp256k1),
_ => None,
}
}
pub const fn a_ptr(self) -> u32 {
match self {
Self::Secp256k1 => K1_A_PTR,
}
}
pub const fn b_ptr(self) -> u32 {
match self {
Self::Secp256k1 => K1_B_PTR,
}
}
pub const fn base_domain(self) -> UintDomain {
match self {
Self::Secp256k1 => UintDomain::K1Base,
}
}
pub fn a_value(self) -> Limbs {
match self {
Self::Secp256k1 => <Secp256k1 as ShortWeierstrassSpec>::A,
}
}
pub fn b_value(self) -> Limbs {
match self {
Self::Secp256k1 => <Secp256k1 as ShortWeierstrassSpec>::B,
}
}
pub const fn scalar_domain(self) -> UintDomain {
match self {
Self::Secp256k1 => UintDomain::K1Scalar,
}
}
pub fn generator(self) -> CurvePoint {
match self {
Self::Secp256k1 => Secp256k1::generator(),
}
}
pub fn endomorphism(self) -> Option<Endomorphism> {
match self {
Self::Secp256k1 => Some(Endomorphism {
beta_ptr: K1_BETA_PTR,
beta: SECP256K1_BETA,
lambda_ptr: K1_LAMBDA_PTR,
lambda: SECP256K1_LAMBDA,
}),
}
}
pub fn point_from_affine(self, x: Limbs, y: Limbs) -> Result<CurvePoint, PrecompileError> {
match self {
Self::Secp256k1 => Secp256k1::point_from_affine(x, y),
}
}
pub fn is_on_curve(self, point: &CurvePoint) -> bool {
match self {
Self::Secp256k1 => Secp256k1::is_on_curve(point),
}
}
pub fn add(self, lhs: CurvePoint, rhs: CurvePoint) -> Result<CurvePoint, PrecompileError> {
match self {
Self::Secp256k1 => Secp256k1::add(lhs, rhs),
}
}
pub fn neg(self, point: CurvePoint) -> Result<CurvePoint, PrecompileError> {
match self {
Self::Secp256k1 => Secp256k1::neg(point),
}
}
pub fn sub(self, lhs: CurvePoint, rhs: CurvePoint) -> Result<CurvePoint, PrecompileError> {
match self {
Self::Secp256k1 => Secp256k1::sub(lhs, rhs),
}
}
pub fn mul_scalar(
self,
point: CurvePoint,
scalar: Limbs,
) -> Result<CurvePoint, PrecompileError> {
match self {
Self::Secp256k1 => Secp256k1::mul_scalar(point, scalar),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum CurveNodeRef {
Value { curve: CurveId, x: Digest, y: Digest },
Add { lhs: Digest, rhs: Digest },
Sub { lhs: Digest, rhs: Digest },
Eq { lhs: Digest, rhs: Digest },
Msm { pairs: Vec<(Digest, Digest)> },
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum CurveBinaryOp {
Add,
Sub,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum CurveOp {
Value(CurveId),
Binary(CurveBinaryOp),
Eq,
Msm,
}
impl CurveOp {
fn decode(args: [Felt; 3]) -> Option<Self> {
match args[0].as_canonical_u64() {
CurvePrecompile::VALUE_OP_ID if args[2] == ZERO => {
let group_ptr = u32::try_from(args[1].as_canonical_u64()).ok()?;
let curve = CurveId::from_group_ptr(group_ptr)?;
Some(Self::Value(curve))
},
CurvePrecompile::ADD_OP_ID if args[1] == ZERO && args[2] == ZERO => {
Some(Self::Binary(CurveBinaryOp::Add))
},
CurvePrecompile::SUB_OP_ID if args[1] == ZERO && args[2] == ZERO => {
Some(Self::Binary(CurveBinaryOp::Sub))
},
CurvePrecompile::EQ_OP_ID if args[1] == ZERO && args[2] == ZERO => Some(Self::Eq),
CurvePrecompile::MSM_OP_ID if args[1] == ZERO && args[2] == ZERO => Some(Self::Msm),
_ => None,
}
}
fn node_type(self) -> NodeType {
match self {
Self::Value(_) | Self::Binary(_) | Self::Eq => NodeType::Join,
Self::Msm => NodeType::PairList,
}
}
}
enum CurveNode {
Value {
curve: CurveId,
lhs: Digest,
rhs: Digest,
},
BinaryOp {
op: CurveBinaryOp,
lhs: Digest,
rhs: Digest,
},
Eq {
lhs: Digest,
rhs: Digest,
},
Msm {
pairs: Vec<(Digest, Digest)>,
},
}
impl CurveNode {
fn parse(op: CurveOp, payload: &Payload) -> Result<Self, PrecompileError> {
Ok(match op {
CurveOp::Value(curve) => {
let (lhs, rhs) = payload.as_join()?;
Self::Value { curve, lhs, rhs }
},
CurveOp::Binary(op) => {
let (lhs, rhs) = payload.as_join()?;
Self::BinaryOp { op, lhs, rhs }
},
CurveOp::Eq => {
let (lhs, rhs) = payload.as_join()?;
Self::Eq { lhs, rhs }
},
CurveOp::Msm => {
let pairs = payload.as_pair_list()?;
Self::Msm { pairs }
},
})
}
}
#[derive(Clone, Copy, Debug, Default)]
pub struct CurvePrecompile;
impl CurvePrecompile {
pub const NAME: &'static str = "curve";
pub const VALUE_OP_ID: u64 = 0;
pub const ADD_OP_ID: u64 = 1;
pub const SUB_OP_ID: u64 = 2;
pub const EQ_OP_ID: u64 = 3;
pub const MSM_OP_ID: u64 = 4;
pub fn id() -> Felt {
precompile_id(Self::NAME)
}
pub fn value_tag(curve: CurveId) -> Tag {
let op_id = Felt::new(Self::VALUE_OP_ID).expect("curve VALUE op id must fit in a felt");
Tag::precompile(Self::id(), [op_id, Felt::from(curve.group_ptr()), ZERO])
.expect("curve precompile id is not framework-reserved")
}
pub fn op_tag(op_id: u64) -> Tag {
let op_id = Felt::new(op_id).expect("curve op id must fit in a felt");
Tag::precompile(Self::id(), [op_id, ZERO, ZERO])
.expect("curve precompile id is not framework-reserved")
}
pub fn msm_tag() -> Tag {
Self::op_tag(Self::MSM_OP_ID)
}
pub fn value_node(curve: CurveId, point: CurvePoint) -> Node {
match point {
CurvePoint::Identity => Self::identity_node(curve),
CurvePoint::Affine { x, y } => Self::affine_node_from_digests(
curve,
UintPrecompile::value_node(curve.base_domain(), x).digest(),
UintPrecompile::value_node(curve.base_domain(), y).digest(),
),
}
}
pub fn identity_node(curve: CurveId) -> Node {
Node::join(Self::value_tag(curve), TRUE_DIGEST, TRUE_DIGEST)
.expect("curve value tag is precompile-owned")
}
pub fn generator_node(curve: CurveId) -> Node {
Self::value_node(curve, curve.generator())
}
pub fn affine_node_from_digests(curve: CurveId, x: Digest, y: Digest) -> Node {
Node::join(Self::value_tag(curve), x, y).expect("curve value tag is precompile-owned")
}
pub fn decode_node(node: &Node) -> Result<Option<CurveNodeRef>, PrecompileError> {
if node.tag().id() != Self::id() {
return Ok(None);
}
let op = CurveOp::decode(node.tag().args()).ok_or(PrecompileError::InvalidNode)?;
let parsed = CurveNode::parse(op, node.payload())?;
Ok(Some(match parsed {
CurveNode::Value { curve, lhs: x, rhs: y } => CurveNodeRef::Value { curve, x, y },
CurveNode::BinaryOp { op: CurveBinaryOp::Add, lhs, rhs } => {
CurveNodeRef::Add { lhs, rhs }
},
CurveNode::BinaryOp { op: CurveBinaryOp::Sub, lhs, rhs } => {
CurveNodeRef::Sub { lhs, rhs }
},
CurveNode::Eq { lhs, rhs } => CurveNodeRef::Eq { lhs, rhs },
CurveNode::Msm { pairs } => CurveNodeRef::Msm { pairs },
}))
}
fn extend_init_nodes_with_point(nodes: &mut Vec<Node>, curve: CurveId, point: CurvePoint) {
if let CurvePoint::Affine { x, y } = point {
let x = UintPrecompile::value_node(curve.base_domain(), x);
let y = UintPrecompile::value_node(curve.base_domain(), y);
nodes.push(x.clone());
nodes.push(y.clone());
nodes.push(Self::affine_node_from_digests(curve, x.digest(), y.digest()));
}
}
fn canonical_value_node(
curve: CurveId,
point: CurvePoint,
context: &mut DeferredContext<'_>,
) -> Result<Node, PrecompileError> {
match point {
CurvePoint::Identity => Ok(Self::identity_node(curve)),
CurvePoint::Affine { x, y } => {
let x = context.register(UintPrecompile::value_node(curve.base_domain(), x))?;
let y = context.register(UintPrecompile::value_node(curve.base_domain(), y))?;
Ok(Self::affine_node_from_digests(curve, x, y))
},
}
}
fn evaluate_msm_term(
expected_curve: Option<CurveId>,
point: Digest,
scalar: Digest,
context: &mut DeferredContext<'_>,
) -> Result<(CurveId, CurvePoint, Limbs), PrecompileError> {
let (point_digest, scalar_digest) = context.evaluate_digest_pair(point, scalar)?;
let point_node = context.get_node(&point_digest).ok_or(PrecompileError::MissingNode)?;
let scalar_node = context.get_node(&scalar_digest).ok_or(PrecompileError::MissingNode)?;
let (point_curve, point) = Self::point_from_value_node(point_node, context)?;
if let Some(expected_curve) = expected_curve
&& expected_curve != point_curve
{
return Err(DeferredError::InvalidPayload.into());
}
let scalar =
UintPrecompile::limbs_from_value_node(scalar_node, point_curve.scalar_domain())?;
Ok((point_curve, point, scalar))
}
fn evaluate_msm(
pairs: &[(Digest, Digest)],
context: &mut DeferredContext<'_>,
) -> Result<(CurveId, CurvePoint), PrecompileError> {
let Some((&(point, scalar), rest)) = pairs.split_first() else {
return Err(DeferredError::InvalidPayload.into());
};
let (curve, point, scalar) = Self::evaluate_msm_term(None, point, scalar, context)?;
if point == CurvePoint::Identity {
return Err(DeferredError::InvalidPayload.into());
}
let mut acc = curve.mul_scalar(point, scalar)?;
for &(point, scalar) in rest {
let (_, point, scalar) = Self::evaluate_msm_term(Some(curve), point, scalar, context)?;
if point == CurvePoint::Identity {
return Err(DeferredError::InvalidPayload.into());
}
let term = curve.mul_scalar(point, scalar)?;
acc = curve.add(acc, term)?;
}
Ok((curve, acc))
}
fn evaluate_point_pair(
context: &mut DeferredContext<'_>,
lhs: Digest,
rhs: Digest,
) -> Result<(CurveId, CurvePoint, CurvePoint), PrecompileError> {
let (lhs, rhs) = context.evaluate_digest_pair(lhs, rhs)?;
let (lhs_curve, lhs) = {
let lhs = context.get_node(&lhs).ok_or(PrecompileError::MissingNode)?;
Self::point_from_value_node(lhs, context)?
};
let (rhs_curve, rhs) = {
let rhs = context.get_node(&rhs).ok_or(PrecompileError::MissingNode)?;
Self::point_from_value_node(rhs, context)?
};
if lhs_curve != rhs_curve {
return Err(DeferredError::InvalidPayload.into());
}
Ok((lhs_curve, lhs, rhs))
}
fn point_from_value_node(
node: &Node,
context: &DeferredContext<'_>,
) -> Result<(CurveId, CurvePoint), PrecompileError> {
let Some(CurveOp::Value(curve)) = CurveOp::decode(node.tag().args()) else {
return Err(DeferredError::InvalidPayload.into());
};
let point = Self::point_of_canonical_node(curve, node, context)?;
Ok((curve, point))
}
fn point_of_canonical_node(
curve: CurveId,
node: &Node,
context: &DeferredContext<'_>,
) -> Result<CurvePoint, PrecompileError> {
let payload = node.payload_for_tag(Self::value_tag(curve))?;
let (x_digest, y_digest) = payload.as_join()?;
Self::point_from_canonical_value_payload(curve, x_digest, y_digest, context)
}
fn point_from_checked_value_payload(
curve: CurveId,
x_digest: Digest,
y_digest: Digest,
context: &DeferredContext<'_>,
) -> Result<CurvePoint, PrecompileError> {
match (x_digest == TRUE_DIGEST, y_digest == TRUE_DIGEST) {
(true, true) => Ok(CurvePoint::Identity),
(true, false) | (false, true) => Err(DeferredError::InvalidPayload.into()),
(false, false) => {
let x_node = context.get_node(&x_digest).ok_or(PrecompileError::MissingNode)?;
let y_node = context.get_node(&y_digest).ok_or(PrecompileError::MissingNode)?;
let x = UintPrecompile::limbs_from_value_node(x_node, curve.base_domain())?;
let y = UintPrecompile::limbs_from_value_node(y_node, curve.base_domain())?;
curve.point_from_affine(x, y)
},
}
}
fn point_from_canonical_value_payload(
curve: CurveId,
x_digest: Digest,
y_digest: Digest,
context: &DeferredContext<'_>,
) -> Result<CurvePoint, PrecompileError> {
match (x_digest == TRUE_DIGEST, y_digest == TRUE_DIGEST) {
(true, true) => Ok(CurvePoint::Identity),
(true, false) | (false, true) => Err(DeferredError::InvalidPayload.into()),
(false, false) => {
let x_node = context.get_node(&x_digest).ok_or(PrecompileError::MissingNode)?;
let y_node = context.get_node(&y_digest).ok_or(PrecompileError::MissingNode)?;
let x = UintPrecompile::limbs_from_value_node(x_node, curve.base_domain())?;
let y = UintPrecompile::limbs_from_value_node(y_node, curve.base_domain())?;
let point = CurvePoint::Affine { x, y };
debug_assert!(curve.is_on_curve(&point));
Ok(point)
},
}
}
}
impl Precompile for CurvePrecompile {
fn name(&self) -> &'static str {
Self::NAME
}
fn id(&self) -> Felt {
Self::id()
}
fn init(&self) -> Vec<Node> {
let mut nodes = Vec::with_capacity(CurveId::ALL.len() * 2);
for curve in CurveId::ALL {
nodes.push(Self::identity_node(curve));
Self::extend_init_nodes_with_point(&mut nodes, curve, curve.generator());
}
nodes
}
fn decode(&self, args: [Felt; 3]) -> Option<NodeType> {
let op = CurveOp::decode(args)?;
Some(op.node_type())
}
fn evaluate(
&self,
args: [Felt; 3],
payload: &Payload,
context: &mut DeferredContext<'_>,
) -> Result<Node, PrecompileError> {
let op = CurveOp::decode(args).ok_or(PrecompileError::InvalidNode)?;
match CurveNode::parse(op, payload)? {
CurveNode::Value { curve, lhs, rhs } => {
match (lhs == TRUE_DIGEST, rhs == TRUE_DIGEST) {
(true, true) => Ok(Self::identity_node(curve)),
(true, false) | (false, true) => Err(DeferredError::InvalidPayload.into()),
(false, false) => {
let (x_digest, y_digest) = context.evaluate_digest_pair(lhs, rhs)?;
let point = Self::point_from_checked_value_payload(
curve, x_digest, y_digest, context,
)?;
Self::canonical_value_node(curve, point, context)
},
}
},
CurveNode::BinaryOp { op, lhs, rhs } => {
let (curve, lhs, rhs) = Self::evaluate_point_pair(context, lhs, rhs)?;
let value = match op {
CurveBinaryOp::Add => curve.add(lhs, rhs)?,
CurveBinaryOp::Sub => curve.sub(lhs, rhs)?,
};
Self::canonical_value_node(curve, value, context)
},
CurveNode::Eq { lhs, rhs } => {
let (_, lhs, rhs) = Self::evaluate_point_pair(context, lhs, rhs)?;
if lhs == rhs {
Ok(Node::TRUE)
} else {
Err(PrecompileError::AssertionFailed)
}
},
CurveNode::Msm { pairs } => {
let (curve, value) = Self::evaluate_msm(&pairs, context)?;
Self::canonical_value_node(curve, value, context)
},
}
}
}
#[cfg(test)]
mod tests {
use alloc::{sync::Arc, vec};
use miden_core::deferred::DeferredState;
use super::*;
use crate::math::{
k1_scalar::K1Scalar,
uint::{UintPrecompile, ZERO_LIMBS},
};
fn state() -> DeferredState {
DeferredState::new(Arc::new(crate::registry())).expect("precompile init must succeed")
}
fn evaluate(state: &mut DeferredState, node: Node) -> Result<Node, PrecompileError> {
let digest = state.register(node)?;
state.require_canonical_node(digest).map(|(_, node)| node.clone())
}
fn assert_invalid_payload<T>(result: Result<T, PrecompileError>) {
let Err(error) = result else {
panic!("expected invalid payload");
};
assert!(
matches!(error.root(), PrecompileError::Other(DeferredError::InvalidPayload)),
"expected invalid payload, got {error:?}",
);
}
fn affine_limbs(point: CurvePoint) -> (Limbs, Limbs) {
match point {
CurvePoint::Affine { x, y } => (x, y),
CurvePoint::Identity => panic!("expected affine point"),
}
}
fn register_affine_point(state: &mut DeferredState, curve: CurveId, point: CurvePoint) -> Node {
let (x, y) = affine_limbs(point);
let x = UintPrecompile::value_node(curve.base_domain(), x);
let y = UintPrecompile::value_node(curve.base_domain(), y);
state.register(x.clone()).expect("x coordinate must register");
state.register(y.clone()).expect("y coordinate must register");
let point = CurvePrecompile::affine_node_from_digests(curve, x.digest(), y.digest());
state.register(point.clone()).expect("point must register");
point
}
#[test]
fn decode_curve_value_tags() {
let precompile = CurvePrecompile;
let curve = CurveId::Secp256k1;
assert_eq!(
CurvePrecompile::value_tag(curve).as_word(),
[
CurvePrecompile::id(),
Felt::from_u32(CurvePrecompile::VALUE_OP_ID as u32),
Felt::from(curve.group_ptr()),
ZERO,
],
);
assert_eq!(
precompile.decode(CurvePrecompile::value_tag(curve).args()),
Some(NodeType::Join)
);
assert_eq!(
precompile.decode([
Felt::from_u32(CurvePrecompile::VALUE_OP_ID as u32),
Felt::from(curve.group_ptr()),
Felt::from_u32(1),
]),
None
);
assert_eq!(
precompile.decode([
Felt::from_u32(CurvePrecompile::VALUE_OP_ID as u32),
Felt::new_unchecked(99),
ZERO,
]),
None
);
}
#[test]
fn decode_curve_operation_tags() {
let precompile = CurvePrecompile;
let curve = CurveId::Secp256k1;
assert_eq!(
precompile.decode(CurvePrecompile::op_tag(CurvePrecompile::ADD_OP_ID).args()),
Some(NodeType::Join)
);
let mut add_with_curve = CurvePrecompile::op_tag(CurvePrecompile::ADD_OP_ID).args();
add_with_curve[1] = Felt::from(curve.group_ptr());
assert_eq!(precompile.decode(add_with_curve), None);
assert_eq!(precompile.decode(CurvePrecompile::op_tag(5).args()), None);
}
#[test]
fn decode_curve_msm_tags() {
let precompile = CurvePrecompile;
let curve = CurveId::Secp256k1;
assert_eq!(
CurvePrecompile::msm_tag().as_word(),
[
CurvePrecompile::id(),
Felt::from_u32(CurvePrecompile::MSM_OP_ID as u32),
ZERO,
ZERO,
],
);
assert_eq!(precompile.decode(CurvePrecompile::msm_tag().args()), Some(NodeType::PairList));
let mut msm_with_curve = CurvePrecompile::msm_tag().args();
msm_with_curve[1] = Felt::from(curve.group_ptr());
assert_eq!(precompile.decode(msm_with_curve), None);
assert_eq!(
precompile.decode([
Felt::from_u32(CurvePrecompile::MSM_OP_ID as u32),
Felt::from(curve.group_ptr()),
Felt::from_u32(1),
]),
None
);
}
#[test]
fn same_curve_add_succeeds() {
let mut state = state();
let curve = CurveId::Secp256k1;
let generator = CurvePrecompile::generator_node(curve);
let identity = CurvePrecompile::identity_node(curve);
let node = Node::join(
CurvePrecompile::op_tag(CurvePrecompile::ADD_OP_ID),
generator.digest(),
identity.digest(),
)
.expect("tag is curve-owned");
assert_eq!(evaluate(&mut state, node).unwrap(), generator);
}
#[test]
fn msm_one_pair_evaluates_point_and_scalar_operands() {
let mut state = state();
let curve = CurveId::Secp256k1;
let generator = CurvePrecompile::generator_node(curve);
let scalar = UintPrecompile::value_node(curve.scalar_domain(), [2, 0, 0, 0, 0, 0, 0, 0]);
state.register(scalar.clone()).expect("scalar must register");
let node = Node::try_pair_list(
CurvePrecompile::msm_tag(),
vec![(generator.digest(), scalar.digest())],
)
.expect("tag is curve-owned");
let expected = CurvePrecompile::value_node(
curve,
curve
.mul_scalar(curve.generator(), [2, 0, 0, 0, 0, 0, 0, 0])
.expect("valid mul_scalar"),
);
assert_eq!(evaluate(&mut state, node).unwrap(), expected);
}
#[test]
fn msm_accumulates_multiple_pairs() {
let mut state = state();
let curve = CurveId::Secp256k1;
let generator = CurvePrecompile::generator_node(curve);
let two_g = curve
.mul_scalar(curve.generator(), [2, 0, 0, 0, 0, 0, 0, 0])
.expect("valid scalar multiplication");
let two_g_node = register_affine_point(&mut state, curve, two_g);
let scalar_2 = UintPrecompile::value_node(curve.scalar_domain(), [2, 0, 0, 0, 0, 0, 0, 0]);
let scalar_3 = UintPrecompile::value_node(curve.scalar_domain(), [3, 0, 0, 0, 0, 0, 0, 0]);
state.register(scalar_2.clone()).expect("scalar must register");
state.register(scalar_3.clone()).expect("scalar must register");
let node = Node::try_pair_list(
CurvePrecompile::msm_tag(),
vec![
(generator.digest(), scalar_2.digest()),
(two_g_node.digest(), scalar_3.digest()),
],
)
.expect("tag is curve-owned");
let two_g_scaled = curve
.mul_scalar(curve.generator(), [2, 0, 0, 0, 0, 0, 0, 0])
.expect("valid scalar multiplication");
let six_g = curve
.mul_scalar(two_g, [3, 0, 0, 0, 0, 0, 0, 0])
.expect("valid scalar multiplication");
let expected = CurvePrecompile::value_node(
curve,
curve.add(two_g_scaled, six_g).expect("valid point addition"),
);
assert_eq!(evaluate(&mut state, node).unwrap(), expected);
}
#[test]
fn msm_accepts_repeated_canonical_base() {
let mut state = state();
let curve = CurveId::Secp256k1;
let generator = CurvePrecompile::generator_node(curve);
let scalar_2 = UintPrecompile::value_node(curve.scalar_domain(), [2, 0, 0, 0, 0, 0, 0, 0]);
let scalar_3 = UintPrecompile::value_node(curve.scalar_domain(), [3, 0, 0, 0, 0, 0, 0, 0]);
state.register(scalar_2.clone()).expect("scalar must register");
state.register(scalar_3.clone()).expect("scalar must register");
let node = Node::try_pair_list(
CurvePrecompile::msm_tag(),
vec![(generator.digest(), scalar_2.digest()), (generator.digest(), scalar_3.digest())],
)
.expect("tag is curve-owned");
let expected = CurvePrecompile::value_node(
curve,
curve
.mul_scalar(curve.generator(), [5, 0, 0, 0, 0, 0, 0, 0])
.expect("valid scalar multiplication"),
);
assert_eq!(evaluate(&mut state, node).unwrap(), expected);
}
#[test]
fn msm_accepts_zero_scalar_terms() {
let mut state = state();
let curve = CurveId::Secp256k1;
let generator = CurvePrecompile::generator_node(curve);
let zero = UintPrecompile::value_node(curve.scalar_domain(), [0; 8]);
state.register(zero.clone()).expect("scalar must register");
let node = Node::try_pair_list(
CurvePrecompile::msm_tag(),
vec![(generator.digest(), zero.digest())],
)
.expect("tag is curve-owned");
let expected = CurvePrecompile::value_node(curve, CurvePoint::Identity);
assert_eq!(evaluate(&mut state, node).unwrap(), expected);
}
#[test]
fn msm_accepts_mixed_zero_and_nonzero_terms() {
let mut state = state();
let curve = CurveId::Secp256k1;
let generator = CurvePrecompile::generator_node(curve);
let two_g = curve
.mul_scalar(curve.generator(), [2, 0, 0, 0, 0, 0, 0, 0])
.expect("valid scalar multiplication");
let two_g_node = register_affine_point(&mut state, curve, two_g);
let zero = UintPrecompile::value_node(curve.scalar_domain(), [0; 8]);
let scalar_3 = UintPrecompile::value_node(curve.scalar_domain(), [3, 0, 0, 0, 0, 0, 0, 0]);
state.register(zero.clone()).expect("scalar must register");
state.register(scalar_3.clone()).expect("scalar must register");
let node = Node::try_pair_list(
CurvePrecompile::msm_tag(),
vec![(generator.digest(), zero.digest()), (two_g_node.digest(), scalar_3.digest())],
)
.expect("tag is curve-owned");
let expected = CurvePrecompile::value_node(
curve,
curve
.mul_scalar(two_g, [3, 0, 0, 0, 0, 0, 0, 0])
.expect("valid scalar multiplication"),
);
assert_eq!(evaluate(&mut state, node).unwrap(), expected);
}
#[test]
fn msm_accepts_multiple_all_zero_terms() {
let mut state = state();
let curve = CurveId::Secp256k1;
let generator = CurvePrecompile::generator_node(curve);
let two_g = curve
.mul_scalar(curve.generator(), [2, 0, 0, 0, 0, 0, 0, 0])
.expect("valid scalar multiplication");
let two_g_node = register_affine_point(&mut state, curve, two_g);
let zero = UintPrecompile::value_node(curve.scalar_domain(), [0; 8]);
state.register(zero.clone()).expect("scalar must register");
let node = Node::try_pair_list(
CurvePrecompile::msm_tag(),
vec![(generator.digest(), zero.digest()), (two_g_node.digest(), zero.digest())],
)
.expect("tag is curve-owned");
let expected = CurvePrecompile::value_node(curve, CurvePoint::Identity);
assert_eq!(evaluate(&mut state, node).unwrap(), expected);
}
#[test]
fn msm_rejects_identity_base_terms() {
let mut state = state();
let curve = CurveId::Secp256k1;
let identity = CurvePrecompile::identity_node(curve);
let scalar = UintPrecompile::value_node(curve.scalar_domain(), [2, 0, 0, 0, 0, 0, 0, 0]);
state.register(identity.clone()).expect("identity must register");
state.register(scalar.clone()).expect("scalar must register");
let node = Node::try_pair_list(
CurvePrecompile::msm_tag(),
vec![(identity.digest(), scalar.digest())],
)
.expect("tag is curve-owned");
assert_invalid_payload(evaluate(&mut state, node));
}
#[test]
fn msm_rejects_wrong_scalar_domain() {
let mut state = state();
let curve = CurveId::Secp256k1;
let generator = CurvePrecompile::generator_node(curve);
let scalar = UintPrecompile::value_node(curve.base_domain(), [2, 0, 0, 0, 0, 0, 0, 0]);
state.register(scalar.clone()).expect("scalar must register");
let node = Node::try_pair_list(
CurvePrecompile::msm_tag(),
vec![(generator.digest(), scalar.digest())],
)
.expect("tag is curve-owned");
assert_invalid_payload(evaluate(&mut state, node));
}
#[test]
fn msm_rejects_non_curve_point_node() {
let mut state = state();
let curve = CurveId::Secp256k1;
let point = UintPrecompile::value_node(curve.base_domain(), [1, 0, 0, 0, 0, 0, 0, 0]);
let scalar = UintPrecompile::value_node(curve.scalar_domain(), [2, 0, 0, 0, 0, 0, 0, 0]);
state.register(point.clone()).expect("point placeholder must register");
state.register(scalar.clone()).expect("scalar must register");
let node = Node::try_pair_list(
CurvePrecompile::msm_tag(),
vec![(point.digest(), scalar.digest())],
)
.expect("tag is curve-owned");
assert_invalid_payload(evaluate(&mut state, node));
}
#[test]
fn mixed_true_payload_is_invalid() {
let identity = CurvePrecompile::identity_node(CurveId::Secp256k1);
let node = CurvePrecompile::affine_node_from_digests(
CurveId::Secp256k1,
TRUE_DIGEST,
identity.digest(),
);
let mut state = state();
assert_invalid_payload(evaluate(&mut state, node));
}
#[test]
fn affine_value_rejects_scalar_field_coordinate_nodes() {
let curve = CurveId::Secp256k1;
let x = UintPrecompile::value_node(UintDomain::K1Scalar, Secp256k1::GENERATOR_X);
let y = UintPrecompile::value_node(UintDomain::K1Scalar, Secp256k1::GENERATOR_Y);
let point = CurvePrecompile::affine_node_from_digests(curve, x.digest(), y.digest());
let mut state = state();
state.register(x).expect("x coordinate must register");
state.register(y).expect("y coordinate must register");
assert_invalid_payload(evaluate(&mut state, point));
}
#[test]
fn off_curve_affine_value_fails_at_registration() {
let curve = CurveId::Secp256k1;
let x = UintPrecompile::value_node(curve.base_domain(), [1, 0, 0, 0, 0, 0, 0, 0]);
let y = UintPrecompile::value_node(curve.base_domain(), [1, 0, 0, 0, 0, 0, 0, 0]);
let point = CurvePrecompile::affine_node_from_digests(curve, x.digest(), y.digest());
let mut state = state();
state.register(x).expect("x coordinate must register");
state.register(y).expect("y coordinate must register");
assert_invalid_payload(state.register(point));
}
#[test]
fn mul_scalar_two_generator_matches_hardcoded_known_answers() {
const K1_2G_X: Limbs = [
0x5c70_9ee5,
0xabac_09b9,
0x8cef_3ca7,
0x5c77_8e4b,
0x95c0_7cd8,
0x3045_406e,
0x41ed_7d6d,
0xc604_7f94,
];
const K1_2G_Y: Limbs = [
0x50cf_e52a,
0x2364_31a9,
0x3266_d0e1,
0xf7f6_3265,
0x466c_eaee,
0xa3c5_8419,
0xa63d_c339,
0x1ae1_68fe,
];
let curve = CurveId::Secp256k1;
let (x, y) =
affine_limbs(curve.mul_scalar(curve.generator(), [2, 0, 0, 0, 0, 0, 0, 0]).unwrap());
assert_eq!((x, y), (K1_2G_X, K1_2G_Y));
}
#[test]
fn fixed_curve_public_pointers_and_generators_validate() {
for curve in CurveId::ALL {
assert_eq!(CurveId::from_group_ptr(curve.group_ptr()), Some(curve));
assert!(curve.is_on_curve(&curve.generator()));
assert!(curve.scalar_domain().is_prime_field());
}
assert_eq!(CurveId::from_group_ptr(99), None);
assert!(!K1Scalar::is_canonical(&K1Scalar::MODULUS));
assert_eq!(
curve_coefficients(),
[
CurveCoefficient {
ptr: CurveId::Secp256k1.a_ptr(),
bound_ptr: CurveId::Secp256k1.base_domain().bound_ptr(),
value: CurveId::Secp256k1.a_value(),
},
CurveCoefficient {
ptr: CurveId::Secp256k1.b_ptr(),
bound_ptr: CurveId::Secp256k1.base_domain().bound_ptr(),
value: CurveId::Secp256k1.b_value(),
},
],
);
}
#[test]
fn mul_scalar_matches_affine_double_and_add_reference() {
fn affine_mul_scalar(curve: CurveId, point: CurvePoint, scalar: Limbs) -> CurvePoint {
let Some(highest_limb) = scalar.iter().rposition(|&limb| limb != 0) else {
return CurvePoint::Identity;
};
let highest_bit =
highest_limb * 32 + (u32::BITS - 1 - scalar[highest_limb].leading_zeros()) as usize;
let mut acc = CurvePoint::Identity;
let mut base = point;
for bit_index in 0..=highest_bit {
if ((scalar[bit_index / 32] >> (bit_index % 32)) & 1) == 1 {
acc = curve.add(acc, base).expect("affine add");
}
if bit_index != highest_bit {
base = curve.add(base, base).expect("affine double");
}
}
acc
}
fn next_u32(state: &mut u64) -> u32 {
*state = state
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
(*state >> 32) as u32
}
let mut state: u64 = 0xdead_beef_cafe_f00d;
for curve in CurveId::ALL {
let modulus = match curve {
CurveId::Secp256k1 => K1Scalar::MODULUS,
};
let minus_one = match curve {
CurveId::Secp256k1 => K1Scalar::minus_one(),
};
let generator = curve.generator();
let generator_128 = affine_mul_scalar(curve, generator, [0, 0, 0, 0, 1, 0, 0, 0]);
let sum = curve.add(generator, generator_128).expect("valid point");
let points = [generator, generator_128, sum];
let edges = [ZERO_LIMBS, [1, 0, 0, 0, 0, 0, 0, 0], [2, 0, 0, 0, 0, 0, 0, 0], minus_one];
for point in points {
for scalar in edges {
assert_eq!(
curve.mul_scalar(point, scalar).expect("projective mul_scalar"),
affine_mul_scalar(curve, point, scalar),
"{curve:?} mul_scalar edge mismatch for scalar {scalar:?}",
);
}
for _ in 0..40 {
let mut scalar = [0u32; 8];
for limb in scalar.iter_mut() {
*limb = next_u32(&mut state);
}
scalar[7] %= modulus[7];
assert!(curve.scalar_domain().is_canonical(&scalar));
assert_eq!(
curve.mul_scalar(point, scalar).expect("projective mul_scalar"),
affine_mul_scalar(curve, point, scalar),
"{curve:?} mul_scalar random mismatch for scalar {scalar:?}",
);
}
}
}
}
}