use crate::utils::{
self, invalid_projective_fallback, Error, HostcallResult, IntoAffineSafe, FAIL_MSG,
};
use alloc::vec::Vec;
use ark_ec::{AffineRepr, CurveConfig};
use ark_ed_on_bls12_381_bandersnatch_ext::CurveHooks;
use sp_runtime_interface::{
pass_by::{PassFatPointerAndRead, PassFatPointerAndWrite},
runtime_interface,
};
pub type BandersnatchConfig = ark_ed_on_bls12_381_bandersnatch_ext::BandersnatchConfig<HostHooks>;
pub type EdwardsConfig = ark_ed_on_bls12_381_bandersnatch_ext::EdwardsConfig<HostHooks>;
pub type EdwardsAffine = ark_ed_on_bls12_381_bandersnatch_ext::EdwardsAffine<HostHooks>;
pub type EdwardsProjective = ark_ed_on_bls12_381_bandersnatch_ext::EdwardsProjective<HostHooks>;
pub type SWConfig = ark_ed_on_bls12_381_bandersnatch_ext::SWConfig<HostHooks>;
pub type SWAffine = ark_ed_on_bls12_381_bandersnatch_ext::SWAffine<HostHooks>;
pub type SWProjective = ark_ed_on_bls12_381_bandersnatch_ext::SWProjective<HostHooks>;
pub type ScalarField = <BandersnatchConfig as CurveConfig>::ScalarField;
#[derive(Copy, Clone)]
pub struct HostHooks;
impl CurveHooks for HostHooks {
fn msm_te(bases: &[EdwardsAffine], scalars: &[ScalarField]) -> EdwardsProjective {
let mut out = utils::buffer_for::<EdwardsAffine>();
match host_calls::ed_on_bls12_381_bandersnatch_msm(
&utils::encode(bases),
&utils::encode(scalars),
&mut out,
) {
Ok(()) => utils::decode::<EdwardsAffine>(&out).expect(FAIL_MSG).into_group(),
Err(Error::DegeneratePoint) => invalid_projective_fallback::<EdwardsConfig>(),
Err(_) => panic!("{FAIL_MSG}"),
}
}
fn mul_projective_te(base: &EdwardsProjective, scalar: &[u64]) -> EdwardsProjective {
let Some(base_aff) = base.into_affine_safe() else {
return invalid_projective_fallback::<EdwardsConfig>();
};
let mut out = utils::buffer_for::<EdwardsAffine>();
match host_calls::ed_on_bls12_381_bandersnatch_mul(
&utils::encode(base_aff),
&utils::encode(scalar),
&mut out,
) {
Ok(()) => utils::decode::<EdwardsAffine>(&out).expect(FAIL_MSG).into_group(),
Err(Error::DegeneratePoint) => invalid_projective_fallback::<EdwardsConfig>(),
Err(_) => panic!("{FAIL_MSG}"),
}
}
}
#[runtime_interface]
pub trait HostCalls {
fn ed_on_bls12_381_bandersnatch_msm(
bases: PassFatPointerAndRead<&[u8]>,
scalars: PassFatPointerAndRead<&[u8]>,
out: PassFatPointerAndWrite<&mut [u8]>,
) -> HostcallResult {
utils::msm_te::<ark_ed_on_bls12_381_bandersnatch::EdwardsConfig>(bases, scalars, out)
}
fn ed_on_bls12_381_bandersnatch_mul(
base: PassFatPointerAndRead<&[u8]>,
scalar: PassFatPointerAndRead<&[u8]>,
out: PassFatPointerAndWrite<&mut [u8]>,
) -> HostcallResult {
utils::mul_te::<ark_ed_on_bls12_381_bandersnatch::EdwardsConfig>(base, scalar, out)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::utils::testing::*;
use ark_ec::{
twisted_edwards::{Affine as TEAffine, Projective as TEProjective, TECurveConfig},
CurveGroup,
};
use ark_ed_on_bls12_381_bandersnatch::{EdwardsConfig as RawConfig, Fq, Fr};
use ark_ff::{AdditiveGroup, MontFp, PrimeField, Zero};
#[test]
fn mul_works() {
mul_te_test::<EdwardsAffine, ark_ed_on_bls12_381_bandersnatch::EdwardsAffine>();
}
#[test]
fn msm_works() {
msm_te_test::<EdwardsAffine, ark_ed_on_bls12_381_bandersnatch::EdwardsAffine>();
}
#[test]
fn mul_works_sw() {
mul_test::<SWAffine, ark_ed_on_bls12_381_bandersnatch::SWAffine>();
}
#[test]
fn msm_works_sw() {
msm_test::<SWAffine, ark_ed_on_bls12_381_bandersnatch::SWAffine>();
}
fn y2_non_subgroup<P: TECurveConfig<BaseField = Fq>>() -> TEAffine<P> {
TEAffine::<P>::get_point_from_y_unchecked(Fq::from(2u64), false)
.expect("y=2 must yield a valid TEAffine point")
}
#[test]
fn host_mul_with_z_zero_result_returns_fallback() {
let proj: TEProjective<RawConfig> = y2_non_subgroup::<RawConfig>().into_group();
let raw_res = <RawConfig as TECurveConfig>::mul_projective(&proj, Fr::MODULUS.0.as_ref());
assert!(raw_res.z.is_zero(), "test precondition: y=2 * Fr::MODULUS must hit z=0");
let scalar_bigint: Vec<u64> = Fr::MODULUS.0.to_vec();
let input_enc = utils::encode(y2_non_subgroup::<EdwardsConfig>());
let scalar_enc = utils::encode(scalar_bigint);
let mut out = utils::buffer_for::<EdwardsAffine>();
let err = host_calls::ed_on_bls12_381_bandersnatch_mul(&input_enc, &scalar_enc, &mut out)
.expect_err("z=0 result must surface as Err(DegeneratePoint)");
assert_eq!(err, Error::DegeneratePoint);
let p_ext: EdwardsProjective = y2_non_subgroup::<EdwardsConfig>().into_group();
let r = <HostHooks as CurveHooks>::mul_projective_te(&p_ext, Fr::MODULUS.0.as_ref());
assert_eq!(
r,
invalid_projective_fallback::<EdwardsConfig>(),
"hook must return all-zero projective on degenerate"
);
}
#[test]
fn mul_projective_with_z_zero_input_returns_fallback() {
use ark_std::{test_rng, UniformRand};
let mut rng = test_rng();
let y = Fq::rand(&mut rng);
let t = Fq::rand(&mut rng);
let p = EdwardsProjective::new_unchecked(Fq::ZERO, y, t, Fq::ZERO);
let r = <HostHooks as CurveHooks>::mul_projective_te(&p, &[7u64, 0, 0, 0]);
assert_eq!(
r,
invalid_projective_fallback::<EdwardsConfig>(),
"z=0 input must yield all-zero coordinate projective"
);
}
#[test]
fn fallback_is_invalid_projective_point() {
let fallback = invalid_projective_fallback::<EdwardsConfig>();
assert!(fallback.x.is_zero(), "fallback x must be zero");
assert!(fallback.y.is_zero(), "fallback y must be zero");
assert!(fallback.t.is_zero(), "fallback t must be zero");
assert!(fallback.z.is_zero(), "fallback z must be zero");
assert!(!fallback.is_zero(), "all-zero projective must NOT be considered identity");
assert!(
fallback.into_affine_safe().is_none(),
"all-zero projective must map to None via IntoAffineSafe",
);
let degenerate = TEProjective::<RawConfig>::new_unchecked(
Fq::ZERO, Fq::from(7u64), Fq::from(11u64), Fq::ZERO, );
assert!(
degenerate.into_affine_safe().is_none(),
"z=0 projective must map to None via IntoAffineSafe",
);
}
fn exceptional_pair() -> (EdwardsAffine, EdwardsAffine, EdwardsAffine) {
let xa: Fq = MontFp!(
"12611587488970178020234800979835231446181428428390492190317266241455236381927"
);
let ya: Fq =
MontFp!("8625363597705895091270672088731506059935752500467284843225771956507605756711");
let xb: Fq =
MontFp!("5253339395048946693631279295832797565125937378490576959411837397991361739535");
let yb: Fq = MontFp!(
"24752777243643877000069062635360441442644758493268974317933177186378585499408"
);
let x_sum: Fq = MontFp!(
"30239213723729448420307207485613680945165091785466061697591732383921178212543"
);
let y_sum: Fq = MontFp!(
"48407687168732128978323921344344221491641898681064657528705691267288289221251"
);
(
EdwardsAffine::new_unchecked(xa, ya),
EdwardsAffine::new_unchecked(xb, yb),
EdwardsAffine::new_unchecked(x_sum, y_sum),
)
}
#[test]
fn hwcd_exceptional_pair_produces_all_zero_projective() {
let (a, b, _expected_sum) = exceptional_pair();
assert!(a.is_on_curve(), "point A must be on curve");
assert!(b.is_on_curve(), "point B must be on curve");
let a_proj: EdwardsProjective = a.into_group();
let b_proj: EdwardsProjective = b.into_group();
let sum = a_proj + b_proj;
assert!(sum.x.is_zero(), "exceptional sum x must be zero");
assert!(sum.y.is_zero(), "exceptional sum y must be zero");
assert!(sum.t.is_zero(), "exceptional sum t must be zero");
assert!(sum.z.is_zero(), "exceptional sum z must be zero");
}
#[test]
fn hwcd_exceptional_pair_recovers_via_sage_sum() {
let (a, b, sage_sum) = exceptional_pair();
let a_proj: EdwardsProjective = a.into_group();
let b_proj: EdwardsProjective = b.into_group();
let sage_sum_proj: EdwardsProjective = sage_sum.into_group();
let ark_sum = a_proj + b_proj;
let a_plus_ark_sum = a_proj + ark_sum;
assert_eq!(
a_plus_ark_sum,
invalid_projective_fallback::<EdwardsConfig>(),
"A + (A+B from arkworks) must produce all-zero projective"
);
let two_a = a_proj + a_proj;
let two_a_plus_b = two_a + b_proj;
let a_plus_sage_sum = a_proj + sage_sum_proj;
assert_eq!(
a_plus_sage_sum.into_affine(),
two_a_plus_b.into_affine(),
"A + (A+B from Sage) must equal 2*A + B"
);
}
#[test]
fn hwcd_exceptional_pair_msm_produces_all_zero() {
use ark_ec::VariableBaseMSM;
let (a, b, _expected_sum) = exceptional_pair();
let bases = vec![a, b];
let scalars = vec![Fr::from(2u64), Fr::from(1u64)];
let result = EdwardsProjective::msm(&bases, &scalars).unwrap();
assert_eq!(
result,
invalid_projective_fallback::<EdwardsConfig>(),
"msm([A, B], [2, 1]) must produce invalid projective fallback"
);
}
#[test]
fn hwcd_exceptional_pair_msm_te_returns_invalid_projective() {
let (a, b, _expected_sum) = exceptional_pair();
let bases = vec![a, b];
let scalars = vec![Fr::from(2u64), Fr::from(1u64)];
let result = <HostHooks as CurveHooks>::msm_te(&bases, &scalars);
assert_eq!(
result,
invalid_projective_fallback::<EdwardsConfig>(),
"msm_te must return invalid projective fallback"
);
}
#[test]
fn y2_point_deserialize_checked_vs_unchecked() {
use ark_scale::ark_serialize::{
CanonicalDeserialize, CanonicalSerialize, Compress, Validate,
};
let p = y2_non_subgroup::<EdwardsConfig>();
assert!(p.is_on_curve(), "y=2 point must be on curve");
assert!(
!p.is_in_correct_subgroup_assuming_on_curve(),
"y=2 point must NOT be in the prime-order subgroup",
);
let mut bytes = Vec::new();
p.serialize_with_mode(&mut bytes, Compress::No).unwrap();
let decoded =
EdwardsAffine::deserialize_with_mode(&bytes[..], Compress::No, Validate::No).unwrap();
assert_eq!(decoded, p);
assert!(
EdwardsAffine::deserialize_with_mode(&bytes[..], Compress::No, Validate::Yes).is_err(),
"Validate::Yes must reject non-subgroup point",
);
}
}