use getrandom::SysRng;
use p256::ecdsa::{Signature, SigningKey, VerifyingKey};
use p256::elliptic_curve::rand_core::UnwrapErr;
use p256::elliptic_curve::sec1::ToSec1Point;
use p256::elliptic_curve::{Field, PrimeField};
use p256::{AffinePoint, FieldBytes, ProjectivePoint, Scalar};
use sha2::{Digest, Sha256};
#[derive(Debug, Clone)]
pub struct AdaptorPreSig {
pub r_prime: AffinePoint,
pub s_prime: Scalar,
}
#[derive(Debug, Clone)]
pub struct WitnessStatement {
pub y_point: AffinePoint,
}
pub fn create_witness() -> (Scalar, WitnessStatement) {
let y = Scalar::random(&mut UnwrapErr(SysRng));
let y_point = (ProjectivePoint::GENERATOR * y).to_affine();
(y, WitnessStatement { y_point })
}
pub fn pre_sign(
signing_key: &SigningKey,
message: &[u8],
witness: &WitnessStatement,
) -> Result<(AdaptorPreSig, Scalar), String> {
let k = Scalar::random(&mut UnwrapErr(SysRng));
let k_point = (ProjectivePoint::GENERATOR * k).to_affine();
let r_prime =
(ProjectivePoint::from(k_point) + ProjectivePoint::from(witness.y_point)).to_affine();
let r = x_coord(&r_prime);
if r == Scalar::ZERO {
return Err("r is zero".into());
}
let e = hash_msg(message);
let x = Option::<Scalar>::from(Scalar::from_repr(p256::FieldBytes::from(
signing_key.to_bytes(),
)))
.unwrap_or(Scalar::ZERO);
let k_inv = invert(&k);
let s_prime = k_inv * (e + r * x);
if s_prime == Scalar::ZERO {
return Err("s' is zero".into());
}
Ok((AdaptorPreSig { r_prime, s_prime }, k))
}
pub fn complete(pre_sig: &AdaptorPreSig, y: &Scalar) -> Result<Signature, String> {
let r = x_coord(&pre_sig.r_prime);
if r == Scalar::ZERO {
return Err("r is zero".into());
}
let y_inv = invert(y);
let s = pre_sig.s_prime * y_inv;
let r_bytes: [u8; 32] = r.to_repr().into();
let s_bytes: [u8; 32] = s.to_repr().into();
Signature::from_scalars(r_bytes, s_bytes).map_err(|e| format!("{e}"))
}
pub fn extract_witness(pre_sig: &AdaptorPreSig, full_sig: &Signature) -> Option<Scalar> {
let (_, s_full) = full_sig.split_scalars();
let s_full_scalar: Scalar = *s_full;
let s_inv = invert(&s_full_scalar);
let y = pre_sig.s_prime * s_inv;
if y == Scalar::ZERO { None } else { Some(y) }
}
pub fn verify_pre_sig(
vk: &VerifyingKey,
message: &[u8],
pre_sig: &AdaptorPreSig,
witness: &WitnessStatement,
) -> bool {
let r = x_coord(&pre_sig.r_prime);
if r == Scalar::ZERO {
return false;
}
let e = hash_msg(message);
let s_inv = invert(&pre_sig.s_prime);
let pk = ProjectivePoint::from(*vk.as_affine());
let u1 = ProjectivePoint::GENERATOR * (e * s_inv);
let u2 = pk * (r * s_inv);
let expected_k_g = u1 + u2;
let k_g = expected_k_g.to_affine();
let expected_r_prime =
(ProjectivePoint::from(k_g) + ProjectivePoint::from(witness.y_point)).to_affine();
expected_r_prime == pre_sig.r_prime
}
fn x_coord(point: &AffinePoint) -> Scalar {
let encoded = point.to_sec1_point(false);
if let Some(x_bytes) = encoded.x() {
let mut arr = [0u8; 32];
arr.copy_from_slice(x_bytes);
let fb = FieldBytes::from(arr);
Option::<Scalar>::from(Scalar::from_repr(fb)).unwrap_or(Scalar::ZERO)
} else {
Scalar::ZERO
}
}
fn hash_msg(msg: &[u8]) -> Scalar {
let mut h = Sha256::new();
h.update(msg);
let fb = FieldBytes::try_from(&h.finalize()[..]).expect("digest is 32 bytes");
Option::<Scalar>::from(Scalar::from_repr(fb)).unwrap_or(Scalar::ZERO)
}
fn invert(s: &Scalar) -> Scalar {
let ct = s.invert();
Option::<Scalar>::from(ct).unwrap_or(Scalar::ZERO)
}
#[cfg(test)]
mod tests {
use super::*;
use p256::elliptic_curve::Generate;
#[test]
fn pre_sign_produces_valid_pre_sig() {
let signing = SigningKey::generate();
let vk = signing.verifying_key();
let (y, witness) = create_witness();
let (pre_sig, _k) = pre_sign(&signing, b"message", &witness).unwrap();
assert!(verify_pre_sig(vk, b"message", &pre_sig, &witness));
let _ = y;
}
#[test]
fn complete_produces_signature() {
let signing = SigningKey::generate();
let (y, witness) = create_witness();
let (pre_sig, _k) = pre_sign(&signing, b"payment", &witness).unwrap();
let full_sig = complete(&pre_sig, &y).unwrap();
let extracted = extract_witness(&pre_sig, &full_sig);
assert!(extracted.is_some());
}
#[test]
fn different_messages_different_pre_sigs() {
let signing = SigningKey::generate();
let (_, w1) = create_witness();
let (_, w2) = create_witness();
let (ps1, _) = pre_sign(&signing, b"msg1", &w1).unwrap();
let (ps2, _) = pre_sign(&signing, b"msg2", &w2).unwrap();
assert_ne!(ps1.r_prime, ps2.r_prime);
}
#[test]
fn wrong_witness_rejected() {
let signing = SigningKey::generate();
let vk = signing.verifying_key();
let (_, w1) = create_witness();
let (_, w2) = create_witness();
let (pre_sig, _) = pre_sign(&signing, b"msg", &w1).unwrap();
assert!(!verify_pre_sig(vk, b"msg", &pre_sig, &w2));
}
#[test]
fn witness_extraction() {
let signing = SigningKey::generate();
let (y, witness) = create_witness();
let (pre_sig, _) = pre_sign(&signing, b"msg", &witness).unwrap();
let full_sig = complete(&pre_sig, &y).unwrap();
let extracted = extract_witness(&pre_sig, &full_sig);
assert!(extracted.is_some());
}
#[test]
fn create_witness_deterministic_point() {
let (y, witness) = create_witness();
let expected = (ProjectivePoint::GENERATOR * y).to_affine();
assert_eq!(witness.y_point, expected);
}
}