use std::fmt::{self, Display};
use std::str::FromStr;
use serde::{Deserialize, Serialize};
use zeroize::{Zeroize, ZeroizeOnDrop};
use crate::encoding::Base64UrlBytes;
use crate::error::{Error, InvalidKeyError, ParseError, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[non_exhaustive]
pub enum OkpCurve {
Ed25519,
Ed448,
X25519,
X448,
}
impl OkpCurve {
pub fn public_key_size(&self) -> usize {
match self {
OkpCurve::Ed25519 | OkpCurve::X25519 => 32,
OkpCurve::Ed448 => 57,
OkpCurve::X448 => 56,
}
}
pub fn private_key_size(&self) -> usize {
match self {
OkpCurve::Ed25519 | OkpCurve::X25519 => 32,
OkpCurve::Ed448 => 57,
OkpCurve::X448 => 56,
}
}
pub fn extended_private_key_size(&self) -> usize {
self.private_key_size() + self.public_key_size()
}
pub fn is_valid_private_key_size(&self, size: usize) -> bool {
size == self.private_key_size() || size == self.extended_private_key_size()
}
pub fn as_str(&self) -> &'static str {
match self {
OkpCurve::Ed25519 => "Ed25519",
OkpCurve::Ed448 => "Ed448",
OkpCurve::X25519 => "X25519",
OkpCurve::X448 => "X448",
}
}
}
impl FromStr for OkpCurve {
type Err = Error;
fn from_str(s: &str) -> Result<Self> {
match s {
"Ed25519" => Ok(OkpCurve::Ed25519),
"Ed448" => Ok(OkpCurve::Ed448),
"X25519" => Ok(OkpCurve::X25519),
"X448" => Ok(OkpCurve::X448),
_ => Err(Error::Parse(ParseError::UnknownCurve(s.to_string()))),
}
}
}
impl Display for OkpCurve {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.as_str())
}
}
#[derive(Clone, PartialEq, Eq, Hash, Serialize, Deserialize, Zeroize, ZeroizeOnDrop)]
#[non_exhaustive]
pub struct OkpParams {
#[zeroize(skip)]
pub crv: OkpCurve,
pub x: Base64UrlBytes,
#[serde(skip_serializing_if = "Option::is_none")]
pub d: Option<Base64UrlBytes>,
}
impl OkpParams {
#[must_use]
pub fn new_public(crv: OkpCurve, x: Base64UrlBytes) -> Self {
Self { crv, x, d: None }
}
#[must_use]
pub fn new_private(crv: OkpCurve, x: Base64UrlBytes, d: Base64UrlBytes) -> Self {
Self { crv, x, d: Some(d) }
}
pub fn is_public_key_only(&self) -> bool {
self.d.is_none()
}
pub fn has_private_key(&self) -> bool {
self.d.is_some()
}
pub fn validate(&self) -> Result<()> {
let expected_public_size = self.crv.public_key_size();
if self.x.len() != expected_public_size {
return Err(InvalidKeyError::InvalidKeySize {
expected: expected_public_size,
actual: self.x.len(),
context: "OKP public key x",
}
.into());
}
if let Some(ref d) = self.d
&& !self.crv.is_valid_private_key_size(d.len())
{
return Err(InvalidKeyError::InvalidKeySize {
expected: self.crv.private_key_size(),
actual: d.len(),
context: "OKP private key d (accepts seed or seed+public format)",
}
.into());
}
Ok(())
}
pub fn private_key_seed(&self) -> Option<&[u8]> {
self.d.as_ref().map(|d| {
let seed_size = self.crv.private_key_size();
if d.len() > seed_size {
&d.as_bytes()[..seed_size]
} else {
d.as_bytes()
}
})
}
#[must_use]
pub fn to_public(&self) -> Self {
Self {
crv: self.crv,
x: self.x.clone(),
d: None,
}
}
}
impl fmt::Debug for OkpParams {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OkpParams")
.field("crv", &self.crv)
.field("x", &format!("[{} bytes]", self.x.len()))
.field("has_private_key", &self.has_private_key())
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_curve_key_sizes() {
assert_eq!(OkpCurve::Ed25519.public_key_size(), 32);
assert_eq!(OkpCurve::Ed25519.private_key_size(), 32);
assert_eq!(OkpCurve::Ed448.public_key_size(), 57);
assert_eq!(OkpCurve::X25519.public_key_size(), 32);
assert_eq!(OkpCurve::X448.public_key_size(), 56);
assert_eq!(OkpCurve::X448.private_key_size(), 56);
}
#[test]
fn test_public_key_only() {
let params = OkpParams::new_public(OkpCurve::Ed25519, Base64UrlBytes::new(vec![0; 32]));
assert!(params.is_public_key_only());
assert!(!params.has_private_key());
}
#[test]
fn test_validate_wrong_size() {
let params = OkpParams::new_public(
OkpCurve::Ed25519,
Base64UrlBytes::new(vec![0; 31]), );
assert!(params.validate().is_err());
}
#[test]
fn test_curve_parsing() {
assert_eq!("Ed25519".parse::<OkpCurve>().unwrap(), OkpCurve::Ed25519);
assert_eq!("Ed448".parse::<OkpCurve>().unwrap(), OkpCurve::Ed448);
assert_eq!("X25519".parse::<OkpCurve>().unwrap(), OkpCurve::X25519);
assert_eq!("X448".parse::<OkpCurve>().unwrap(), OkpCurve::X448);
assert!("unknown".parse::<OkpCurve>().is_err());
}
#[test]
fn test_json_roundtrip() {
let original = OkpParams::new_public(OkpCurve::Ed25519, Base64UrlBytes::new(vec![1; 32]));
let json = serde_json::to_string(&original).unwrap();
let decoded: OkpParams = serde_json::from_str(&json).unwrap();
assert_eq!(original, decoded);
}
#[test]
fn test_ed448_extended_private_key() {
assert_eq!(OkpCurve::Ed448.private_key_size(), 57);
assert_eq!(OkpCurve::Ed448.extended_private_key_size(), 114);
assert!(OkpCurve::Ed448.is_valid_private_key_size(57));
assert!(OkpCurve::Ed448.is_valid_private_key_size(114));
assert!(!OkpCurve::Ed448.is_valid_private_key_size(32));
assert!(!OkpCurve::Ed448.is_valid_private_key_size(100));
}
#[test]
fn test_ed448_seed_format_validates() {
let params = OkpParams::new_private(
OkpCurve::Ed448,
Base64UrlBytes::new(vec![0; 57]),
Base64UrlBytes::new(vec![1; 57]),
);
assert!(params.validate().is_ok());
}
#[test]
fn test_ed448_extended_format_validates() {
let params = OkpParams::new_private(
OkpCurve::Ed448,
Base64UrlBytes::new(vec![0; 57]),
Base64UrlBytes::new(vec![1; 114]),
);
assert!(params.validate().is_ok());
}
#[test]
fn test_private_key_seed_extraction() {
let seed = vec![1u8; 57];
let public = vec![2u8; 57];
let mut extended = seed.clone();
extended.extend_from_slice(&public);
let params = OkpParams::new_private(
OkpCurve::Ed448,
Base64UrlBytes::new(vec![0; 57]),
Base64UrlBytes::new(extended),
);
let extracted_seed = params.private_key_seed().unwrap();
assert_eq!(extracted_seed, &seed[..]);
}
#[test]
fn test_ed25519_extended_key_sizes() {
assert_eq!(OkpCurve::Ed25519.private_key_size(), 32);
assert_eq!(OkpCurve::Ed25519.extended_private_key_size(), 64);
assert!(OkpCurve::Ed25519.is_valid_private_key_size(32));
assert!(OkpCurve::Ed25519.is_valid_private_key_size(64));
}
#[test]
fn test_x448_sizes_and_validation() {
assert_eq!(OkpCurve::X448.public_key_size(), 56);
assert_eq!(OkpCurve::X448.private_key_size(), 56);
assert_eq!(OkpCurve::X448.extended_private_key_size(), 112);
assert!(OkpCurve::X448.is_valid_private_key_size(56));
assert!(OkpCurve::X448.is_valid_private_key_size(112));
let params = OkpParams::new_private(
OkpCurve::X448,
Base64UrlBytes::new(vec![0; 56]),
Base64UrlBytes::new(vec![1; 56]),
);
assert!(params.validate().is_ok());
}
}