use crate::error::VerificationFailure;
use crate::schnorr::traits::{Prover, Randomized, Verifier};
use elliptic_curve::group::Curve;
use elliptic_curve::ops::MulByGenerator;
use elliptic_curve::point::PointCompression;
use elliptic_curve::rand_core::CryptoRngCore;
use elliptic_curve::sec1::{EncodedPoint, FromEncodedPoint, ModulusSize, ToEncodedPoint};
use elliptic_curve::{
CurveArithmetic, Error, FieldBytes, FieldBytesSize, NonZeroScalar, PrimeCurve, PrimeField,
PublicKey, ScalarPrimitive, SecretKey,
};
use zeroize::{Zeroize, ZeroizeOnDrop};
pub struct SchnorrECProver<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
secret: NonZeroScalar<C>,
nonce: NonZeroScalar<C>,
}
impl<C> Prover for SchnorrECProver<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
type SecretKey = SecretKey<C>;
type Commitment = Commitment<C>;
type Nonce = Nonce<C>;
type Challenge = Challenge<C>;
type Answer = Answer<C>;
fn new_with_nonce(secret_key: &Self::SecretKey, nonce: Self::Nonce) -> Self {
SchnorrECProver {
secret: secret_key.into(),
nonce: nonce.into(),
}
}
fn nonce(&self) -> Self::Nonce {
self.nonce.into()
}
fn commitment(&self) -> Self::Commitment {
Commitment::new(C::ProjectivePoint::mul_by_generator(&self.nonce).to_affine())
}
fn answer(self, challenge: Self::Challenge) -> Self::Answer {
Answer::new(*self.nonce + *challenge.as_scalar() * *self.secret)
}
}
impl<C> Zeroize for SchnorrECProver<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
fn zeroize(&mut self) {
self.secret.zeroize();
self.nonce.zeroize();
}
}
impl<C> Drop for SchnorrECProver<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
fn drop(&mut self) {
self.zeroize();
}
}
impl<C> ZeroizeOnDrop for SchnorrECProver<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
}
pub struct SchnorrECVerifier<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
public_key: C::ProjectivePoint,
commitment: C::ProjectivePoint,
challenge: Challenge<C>,
}
impl<C> Verifier for SchnorrECVerifier<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
type PublicKey = PublicKey<C>;
type Commitment = Commitment<C>;
type Challenge = Challenge<C>;
type Answer = Answer<C>;
fn new_with_challenge(
public_key: &Self::PublicKey,
commitment: Self::Commitment,
challenge: Self::Challenge,
) -> Self {
SchnorrECVerifier {
public_key: public_key.to_projective(),
commitment: commitment.to_affine().into(),
challenge,
}
}
fn challenge(&self) -> Self::Challenge {
self.challenge
}
fn verify(self, answer: Self::Answer) -> Result<(), VerificationFailure> {
let point1 = self.commitment + self.public_key * self.challenge.as_scalar();
let point2 = C::ProjectivePoint::mul_by_generator(answer.as_scalar());
if point1 == point2 {
Ok(())
} else {
Err(VerificationFailure)
}
}
}
pub struct Nonce<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
scalar: NonZeroScalar<C>,
}
impl<C> Nonce<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
pub fn new(scalar: NonZeroScalar<C>) -> Self {
Self { scalar }
}
pub fn to_nonzero_scalar(&self) -> NonZeroScalar<C> {
self.scalar
}
}
impl<C> Randomized for Nonce<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
fn random(rng: &mut impl CryptoRngCore) -> Self {
Self::new(NonZeroScalar::random(rng))
}
}
impl<C> From<NonZeroScalar<C>> for Nonce<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
fn from(scalar: NonZeroScalar<C>) -> Self {
Self::new(scalar)
}
}
impl<C> From<Nonce<C>> for NonZeroScalar<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
fn from(nonce: Nonce<C>) -> Self {
nonce.to_nonzero_scalar()
}
}
impl<C> Zeroize for Nonce<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
fn zeroize(&mut self) {
self.scalar.zeroize();
}
}
impl<C> Drop for Nonce<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
fn drop(&mut self) {
self.zeroize();
}
}
pub struct Commitment<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
point: C::AffinePoint,
}
impl<C> Commitment<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
pub fn new(point: C::AffinePoint) -> Self {
Self { point }
}
pub fn from_sec1_bytes(bytes: &[u8]) -> Result<Self, Error> {
let encoded_point = EncodedPoint::<C>::from_bytes(bytes)?;
Self::from_encoded_point(&encoded_point)
}
pub fn from_encoded_point(encoded_point: &EncodedPoint<C>) -> Result<Self, Error> {
Option::from(C::AffinePoint::from_encoded_point(encoded_point).map(Self::new)).ok_or(Error)
}
pub fn to_affine(&self) -> C::AffinePoint {
self.point
}
pub fn to_sec1_bytes(&self) -> Vec<u8>
where
C: PointCompression,
{
self.point
.to_encoded_point(C::COMPRESS_POINTS)
.as_bytes()
.to_vec()
}
pub fn to_encoded_point(&self, compress: bool) -> EncodedPoint<C> {
self.point.to_encoded_point(compress)
}
}
#[derive(Copy, Clone)]
pub struct Challenge<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
scalar: C::Scalar,
}
impl<C> Challenge<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
pub fn from_bytes(bytes: &FieldBytes<C>) -> Result<Self, Error> {
Option::from(ScalarPrimitive::from_bytes(bytes).map(|scalar| Self {
scalar: scalar.into(),
}))
.ok_or(Error)
}
pub fn from_slice(bytes: &[u8]) -> Result<Self, Error> {
ScalarPrimitive::from_slice(bytes).map(|scalar| Self {
scalar: scalar.into(),
})
}
pub fn as_scalar(&self) -> &C::Scalar {
&self.scalar
}
pub fn to_bytes(&self) -> FieldBytes<C> {
self.scalar.to_repr()
}
}
impl<C> Randomized for Challenge<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
fn random(rng: &mut impl CryptoRngCore) -> Self {
Self {
scalar: ScalarPrimitive::from(NonZeroScalar::<C>::random(rng)).into(),
}
}
}
pub struct Answer<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
scalar: C::Scalar,
}
impl<C> Answer<C>
where
C: CurveArithmetic + PrimeCurve,
FieldBytesSize<C>: ModulusSize,
C::AffinePoint: FromEncodedPoint<C> + ToEncodedPoint<C>,
{
pub fn new(scalar: C::Scalar) -> Self {
Self { scalar }
}
pub fn from_bytes(bytes: &FieldBytes<C>) -> Result<Self, Error> {
Option::from(ScalarPrimitive::from_bytes(bytes).map(|scalar| Self {
scalar: scalar.into(),
}))
.ok_or(Error)
}
pub fn from_slice(bytes: &[u8]) -> Result<Self, Error> {
ScalarPrimitive::from_slice(bytes)
.map(|scalar| Self {
scalar: scalar.into(),
})
.map_err(|_| Error)
}
pub fn as_scalar(&self) -> &C::Scalar {
&self.scalar
}
pub fn to_bytes(&self) -> FieldBytes<C> {
self.scalar.to_repr()
}
}
#[cfg(test)]
mod tests {
use super::*;
use p256::NistP256;
use rand::rngs::OsRng;
use std::sync::mpsc::channel;
use std::thread;
#[test]
fn it_works() {
let secret_key = SecretKey::<NistP256>::random(&mut OsRng);
let public_key = secret_key.public_key();
let (client_sender, server_receiver) = channel::<String>();
let (server_sender, client_receiver) = channel::<String>();
let client_thread = thread::spawn(move || {
let prover = SchnorrECProver::new(&secret_key, &mut OsRng);
let commitment = prover.commitment();
let commitment = hex::encode(commitment.to_sec1_bytes());
client_sender.send(commitment).unwrap();
let challenge = client_receiver.recv().unwrap();
let challenge = hex::decode(challenge).unwrap();
let answer = prover.answer(Challenge::from_slice(&challenge).unwrap());
let answer = hex::encode(answer.to_bytes());
client_sender.send(answer).unwrap();
});
let server_thread = thread::spawn(move || {
let commitment = server_receiver.recv().unwrap();
let commitment = hex::decode(commitment).unwrap();
let verifier = SchnorrECVerifier::new(
&public_key,
Commitment::from_sec1_bytes(&commitment).unwrap(),
&mut OsRng,
);
let challenge = verifier.challenge();
let challenge = hex::encode(challenge.to_bytes());
server_sender.send(challenge).unwrap();
let answer = server_receiver.recv().unwrap();
let answer = hex::decode(answer).unwrap();
verifier
.verify(Answer::from_slice(&answer).unwrap())
.unwrap();
});
client_thread.join().unwrap();
server_thread.join().unwrap();
}
}