use core::marker::PhantomData;
#[cfg(feature = "alloc")]
use alloc::vec::Vec;
use crate::common::error::{CryptoError, Result};
use crate::kes::hash::KesHashAlgorithm;
use crate::kes::{KesAlgorithm, KesError, Period};
#[derive(Debug)]
pub struct SumKes<D, H>(PhantomData<(D, H)>)
where
D: KesAlgorithm,
H: KesHashAlgorithm;
#[derive(Debug)]
pub struct SumSigningKey<D, H>
where
D: KesAlgorithm,
H: KesHashAlgorithm,
{
pub(crate) sk: D::SigningKey,
pub(crate) r1_seed: Option<Vec<u8>>,
pub(crate) vk0: D::VerificationKey,
pub(crate) vk1: D::VerificationKey,
_phantom: PhantomData<H>,
}
#[derive(Clone, PartialEq, Eq)]
pub struct SumSignature<D, H>
where
D: KesAlgorithm,
H: KesHashAlgorithm,
{
pub(crate) sigma: D::Signature,
pub(crate) vk0: D::VerificationKey,
pub(crate) vk1: D::VerificationKey,
_phantom: PhantomData<H>,
}
impl<D, H> core::fmt::Debug for SumSignature<D, H>
where
D: KesAlgorithm,
D::Signature: core::fmt::Debug,
D::VerificationKey: core::fmt::Debug,
H: KesHashAlgorithm,
{
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("SumSignature")
.field("sigma", &self.sigma)
.field("vk0", &"<VK>")
.field("vk1", &"<VK>")
.finish()
}
}
impl<D, H> KesAlgorithm for SumKes<D, H>
where
D: KesAlgorithm,
D::VerificationKey: Clone,
H: KesHashAlgorithm,
{
type VerificationKey = Vec<u8>; type SigningKey = SumSigningKey<D, H>;
type Signature = SumSignature<D, H>;
type Context = D::Context;
const ALGORITHM_NAME: &'static str = D::ALGORITHM_NAME;
const SEED_SIZE: usize = D::SEED_SIZE;
const VERIFICATION_KEY_SIZE: usize = H::OUTPUT_SIZE;
const SIGNING_KEY_SIZE: usize =
D::SIGNING_KEY_SIZE + D::SEED_SIZE + 2 * D::VERIFICATION_KEY_SIZE;
const SIGNATURE_SIZE: usize = D::SIGNATURE_SIZE + 2 * D::VERIFICATION_KEY_SIZE;
fn total_periods() -> Period {
2 * D::total_periods()
}
fn derive_verification_key(signing_key: &Self::SigningKey) -> Result<Self::VerificationKey> {
let vk0_bytes = D::raw_serialize_verification_key_kes(&signing_key.vk0);
let vk1_bytes = D::raw_serialize_verification_key_kes(&signing_key.vk1);
Ok(H::hash_concat(&vk0_bytes, &vk1_bytes))
}
fn sign_kes(
context: &Self::Context,
period: Period,
message: &[u8],
signing_key: &Self::SigningKey,
) -> Result<Self::Signature> {
let t_half = D::total_periods();
let sigma = if period < t_half {
D::sign_kes(context, period, message, &signing_key.sk)?
} else {
D::sign_kes(context, period - t_half, message, &signing_key.sk)?
};
Ok(SumSignature {
sigma,
vk0: signing_key.vk0.clone(),
vk1: signing_key.vk1.clone(),
_phantom: PhantomData,
})
}
fn verify_kes(
context: &Self::Context,
verification_key: &Self::VerificationKey,
period: Period,
message: &[u8],
signature: &Self::Signature,
) -> Result<()> {
let vk0_bytes = D::raw_serialize_verification_key_kes(&signature.vk0);
let vk1_bytes = D::raw_serialize_verification_key_kes(&signature.vk1);
let computed_vk = H::hash_concat(&vk0_bytes, &vk1_bytes);
if &computed_vk != verification_key {
return Err(CryptoError::KesError(KesError::VerificationFailed));
}
let t_half = D::total_periods();
if period < t_half {
D::verify_kes(context, &signature.vk0, period, message, &signature.sigma)
} else {
D::verify_kes(
context,
&signature.vk1,
period - t_half,
message,
&signature.sigma,
)
}
}
fn update_kes(
context: &Self::Context,
mut signing_key: Self::SigningKey,
period: Period,
) -> Result<Option<Self::SigningKey>> {
let t_half = D::total_periods();
if period + 1 >= 2 * t_half {
D::forget_signing_key_kes(signing_key.sk);
return Ok(None);
}
if period + 1 == t_half {
let r1_seed = signing_key
.r1_seed
.take()
.ok_or(CryptoError::KesError(KesError::KeyExpired))?;
let sk1 = D::gen_key_kes_from_seed_bytes(&r1_seed)?;
D::forget_signing_key_kes(signing_key.sk);
Ok(Some(SumSigningKey {
sk: sk1,
r1_seed: None, vk0: signing_key.vk0,
vk1: signing_key.vk1,
_phantom: PhantomData,
}))
} else if period + 1 < t_half {
let updated_sk = D::update_kes(context, signing_key.sk, period)?;
match updated_sk {
Some(sk) => Ok(Some(SumSigningKey {
sk,
r1_seed: signing_key.r1_seed,
vk0: signing_key.vk0,
vk1: signing_key.vk1,
_phantom: PhantomData,
})),
None => Ok(None),
}
} else {
let adjusted_period = period - t_half;
let updated_sk = D::update_kes(context, signing_key.sk, adjusted_period)?;
match updated_sk {
Some(sk) => Ok(Some(SumSigningKey {
sk,
r1_seed: None,
vk0: signing_key.vk0,
vk1: signing_key.vk1,
_phantom: PhantomData,
})),
None => Ok(None),
}
}
}
fn gen_key_kes_from_seed_bytes(seed: &[u8]) -> Result<Self::SigningKey> {
if seed.len() != Self::SEED_SIZE {
return Err(CryptoError::KesError(KesError::InvalidSeedLength {
expected: Self::SEED_SIZE,
actual: seed.len(),
}));
}
let (r0_bytes, r1_bytes) = H::expand_seed(seed);
let sk0 = D::gen_key_kes_from_seed_bytes(&r0_bytes)?;
let vk0 = D::derive_verification_key(&sk0)?;
let sk1 = D::gen_key_kes_from_seed_bytes(&r1_bytes)?;
let vk1 = D::derive_verification_key(&sk1)?;
D::forget_signing_key_kes(sk1);
Ok(SumSigningKey {
sk: sk0,
r1_seed: Some(r1_bytes),
vk0,
vk1,
_phantom: PhantomData,
})
}
#[cfg(feature = "alloc")]
fn raw_serialize_verification_key_kes(key: &Self::VerificationKey) -> Vec<u8> {
key.clone()
}
fn raw_deserialize_verification_key_kes(bytes: &[u8]) -> Option<Self::VerificationKey> {
if bytes.len() == Self::VERIFICATION_KEY_SIZE {
Some(bytes.to_vec())
} else {
None
}
}
#[cfg(feature = "alloc")]
fn raw_serialize_signature_kes(signature: &Self::Signature) -> Vec<u8> {
let mut result = D::raw_serialize_signature_kes(&signature.sigma);
result.extend_from_slice(&D::raw_serialize_verification_key_kes(&signature.vk0));
result.extend_from_slice(&D::raw_serialize_verification_key_kes(&signature.vk1));
result
}
fn raw_deserialize_signature_kes(bytes: &[u8]) -> Option<Self::Signature> {
if bytes.len() != Self::SIGNATURE_SIZE {
return None;
}
let sig_bytes = &bytes[0..D::SIGNATURE_SIZE];
let vk0_offset = D::SIGNATURE_SIZE;
let vk1_offset = vk0_offset + D::VERIFICATION_KEY_SIZE;
let sigma = D::raw_deserialize_signature_kes(sig_bytes)?;
let vk0 = D::raw_deserialize_verification_key_kes(&bytes[vk0_offset..vk1_offset])?;
let vk1 = D::raw_deserialize_verification_key_kes(&bytes[vk1_offset..])?;
Some(SumSignature {
sigma,
vk0,
vk1,
_phantom: PhantomData,
})
}
fn forget_signing_key_kes(signing_key: Self::SigningKey) {
D::forget_signing_key_kes(signing_key.sk);
}
}
use crate::dsign::ed25519::Ed25519;
use crate::kes::hash::Blake2b256;
use crate::kes::single::SingleKes;
pub type Sum0Kes = SingleKes<Ed25519>;
pub type Sum1Kes = SumKes<Sum0Kes, Blake2b256>;
pub type Sum2Kes = SumKes<Sum1Kes, Blake2b256>;
pub type Sum3Kes = SumKes<Sum2Kes, Blake2b256>;
pub type Sum4Kes = SumKes<Sum3Kes, Blake2b256>;
pub type Sum5Kes = SumKes<Sum4Kes, Blake2b256>;
pub type Sum6Kes = SumKes<Sum5Kes, Blake2b256>;
pub type Sum7Kes = SumKes<Sum6Kes, Blake2b256>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sum1_total_periods() {
assert_eq!(Sum1Kes::total_periods(), 2);
}
#[test]
fn sum2_total_periods() {
assert_eq!(Sum2Kes::total_periods(), 4);
}
#[test]
fn sum3_total_periods() {
assert_eq!(Sum3Kes::total_periods(), 8);
}
#[test]
fn sum4_total_periods() {
assert_eq!(Sum4Kes::total_periods(), 16);
}
#[test]
fn sum1_key_generation_and_derivation() {
let seed = vec![1u8; Sum1Kes::SEED_SIZE];
let sk = Sum1Kes::gen_key_kes_from_seed_bytes(&seed).unwrap();
let vk = Sum1Kes::derive_verification_key(&sk).unwrap();
assert_eq!(vk.len(), Sum1Kes::VERIFICATION_KEY_SIZE);
assert_eq!(vk.len(), 32); }
#[test]
fn sum1_sign_and_verify_period_0() {
let seed = vec![2u8; Sum1Kes::SEED_SIZE];
let sk = Sum1Kes::gen_key_kes_from_seed_bytes(&seed).unwrap();
let vk = Sum1Kes::derive_verification_key(&sk).unwrap();
let msg = b"period-0-message";
let sig = Sum1Kes::sign_kes(&(), 0, msg, &sk).unwrap();
Sum1Kes::verify_kes(&(), &vk, 0, msg, &sig).unwrap();
}
#[test]
fn sum1_sign_and_verify_period_1() {
let seed = vec![3u8; Sum1Kes::SEED_SIZE];
let sk = Sum1Kes::gen_key_kes_from_seed_bytes(&seed).unwrap();
let vk = Sum1Kes::derive_verification_key(&sk).unwrap();
let sk = Sum1Kes::update_kes(&(), sk, 0).unwrap().unwrap();
let msg = b"period-1-message";
let sig = Sum1Kes::sign_kes(&(), 1, msg, &sk).unwrap();
Sum1Kes::verify_kes(&(), &vk, 1, msg, &sig).unwrap();
}
#[test]
fn sum1_key_expires_after_period_1() {
let seed = vec![4u8; Sum1Kes::SEED_SIZE];
let sk = Sum1Kes::gen_key_kes_from_seed_bytes(&seed).unwrap();
let sk = Sum1Kes::update_kes(&(), sk, 0).unwrap().unwrap();
let updated = Sum1Kes::update_kes(&(), sk, 1).unwrap();
assert!(updated.is_none(), "Sum1Kes should expire after period 1");
}
#[test]
fn sum2_full_lifecycle() {
let seed = vec![5u8; Sum2Kes::SEED_SIZE];
let mut sk = Sum2Kes::gen_key_kes_from_seed_bytes(&seed).unwrap();
let vk = Sum2Kes::derive_verification_key(&sk).unwrap();
for period in 0..4 {
let msg = format!("period-{}", period);
let sig = Sum2Kes::sign_kes(&(), period, msg.as_bytes(), &sk).unwrap();
Sum2Kes::verify_kes(&(), &vk, period, msg.as_bytes(), &sig).unwrap();
if period < 3 {
sk = Sum2Kes::update_kes(&(), sk, period).unwrap().unwrap();
}
}
let updated = Sum2Kes::update_kes(&(), sk, 3).unwrap();
assert!(updated.is_none(), "Sum2Kes should expire after period 3");
}
#[test]
fn sum1_signature_serialization() {
let seed = vec![6u8; Sum1Kes::SEED_SIZE];
let sk = Sum1Kes::gen_key_kes_from_seed_bytes(&seed).unwrap();
let vk = Sum1Kes::derive_verification_key(&sk).unwrap();
let msg = b"serialize-test";
let sig = Sum1Kes::sign_kes(&(), 0, msg, &sk).unwrap();
let sig_bytes = Sum1Kes::raw_serialize_signature_kes(&sig);
let sig_restored = Sum1Kes::raw_deserialize_signature_kes(&sig_bytes).unwrap();
Sum1Kes::verify_kes(&(), &vk, 0, msg, &sig_restored).unwrap();
}
}