use crate::error::{Error, Result};
use constant_time_eq::constant_time_eq;
use k256::elliptic_curve::group::prime::PrimeCurveAffine;
use k256::elliptic_curve::sec1::ToEncodedPoint;
use k256::{AffinePoint, NonZeroScalar};
const SIZE: usize = 65;
#[derive(Debug, Clone, Copy)]
pub struct PublicKey {
bytes: [u8; SIZE],
}
impl PartialEq for PublicKey {
fn eq(&self, other: &Self) -> bool {
constant_time_eq(&self.bytes, &other.bytes)
}
}
impl Eq for PublicKey {}
impl PublicKey {
pub fn from_bytes(bytes: [u8; SIZE]) -> Self {
Self { bytes }
}
pub fn from_slice(slice: &[u8]) -> Result<Self> {
if slice.len() != SIZE {
Err(Error::LengthError(format!(
"PublicKey size must be {} bytes got {} bytes.",
SIZE,
slice.len()
)))
} else {
let mut new_bytes = [0u8; SIZE];
new_bytes.copy_from_slice(slice);
Ok(Self::from_bytes(new_bytes))
}
}
pub fn from_hex(hex: &str) -> Result<Self> {
if hex.len() != SIZE * 2 {
Err(Error::LengthError(format!(
"PublicKey hex must be {} characters got {} characters.",
SIZE * 2,
hex.len()
)))
} else {
let bytes = hex::decode(hex)?;
Self::from_slice(&bytes)
}
}
pub fn as_bytes(&self) -> [u8; SIZE] {
self.bytes
}
pub fn as_compressed_bytes(&self) -> [u8; 33] {
let pk = k256::PublicKey::from_sec1_bytes(&self.bytes).unwrap();
let bytes = pk.to_encoded_point(true);
bytes.as_bytes().try_into().unwrap()
}
pub fn as_slice(&self) -> &[u8] {
&self.bytes
}
pub fn as_hex(&self) -> String {
hex::encode(&self.bytes)
}
pub fn derive_child(&self, other: [u8; 32]) -> Result<Self> {
let current = k256::PublicKey::from_sec1_bytes(&self.bytes).unwrap();
let child_scalar = Option::<NonZeroScalar>::from(NonZeroScalar::from_repr(other.into()))
.ok_or(Error::InvalidPublicKey)?;
let child_point = current.to_projective() + (AffinePoint::generator() * *child_scalar);
let derived = k256::PublicKey::from_affine(child_point.into())
.map_err(|_| Error::InvalidPublicKey)?;
let bytes = derived.to_encoded_point(false);
Self::from_slice(bytes.as_bytes())
}
}
#[cfg(test)]
mod tests {
use crate::PrivateKey;
use super::*;
#[test]
fn test_public_key() {
let key = PrivateKey::new();
let pk = key.public_key();
let pk_bytes = pk.as_bytes();
let pk_hex = pk.as_hex();
let pk_slice = pk.as_slice();
let pk_compressed = pk.as_compressed_bytes();
assert_eq!(pk_bytes.len(), SIZE);
assert_eq!(pk_hex.len(), SIZE * 2);
assert_eq!(pk_slice.len(), SIZE);
assert_eq!(pk_compressed.len(), 33);
assert_eq!(pk, PublicKey::from_slice(pk_slice).unwrap());
assert_eq!(pk, PublicKey::from_hex(&pk_hex).unwrap());
}
}