use std::fmt;
#[cfg(feature = "serde")]
use serde::{Deserialize, Serialize};
use zeroize::{Zeroize, ZeroizeOnDrop};
use crate::classic::crypto_kdf::{crypto_kdf_derive_from_key, validate_subkey_length};
use crate::constants::{CRYPTO_KDF_CONTEXTBYTES, CRYPTO_KDF_KEYBYTES};
use crate::error::Error;
use crate::types::*;
pub type Key = StackByteArray<CRYPTO_KDF_KEYBYTES>;
pub type Context = StackByteArray<CRYPTO_KDF_CONTEXTBYTES>;
#[cfg_attr(feature = "serde", derive(Zeroize, Clone, Serialize, Deserialize))]
#[cfg_attr(not(feature = "serde"), derive(Zeroize, Clone))]
pub struct Kdf<
Key: ByteArray<CRYPTO_KDF_KEYBYTES> + Zeroize + ZeroizeOnDrop,
Context: ByteArray<CRYPTO_KDF_CONTEXTBYTES> + Zeroize,
> {
main_key: Key,
context: Context,
}
impl<
Key: ByteArray<CRYPTO_KDF_KEYBYTES> + Zeroize + ZeroizeOnDrop,
Context: ByteArray<CRYPTO_KDF_CONTEXTBYTES> + Zeroize,
> fmt::Debug for Kdf<Key, Context>
{
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Kdf")
.field("main_key", &"[REDACTED]")
.field("context", &self.context.as_slice())
.finish()
}
}
pub type StackKdf = Kdf<Key, Context>;
#[cfg(any(all(feature = "protected", any(unix, windows)), all(doc, not(doctest))))]
#[cfg_attr(all(feature = "nightly", doc), doc(cfg(feature = "protected")))]
pub mod protected {
use super::*;
pub use crate::protected::*;
pub type Key = HeapByteArray<CRYPTO_KDF_KEYBYTES>;
pub type Context = HeapByteArray<CRYPTO_KDF_CONTEXTBYTES>;
pub type LockedKdf = Kdf<Locked<Key>, Locked<Context>>;
}
impl<
Key: NewByteArray<CRYPTO_KDF_KEYBYTES> + Zeroize + ZeroizeOnDrop,
Context: NewByteArray<CRYPTO_KDF_CONTEXTBYTES> + Zeroize,
> Kdf<Key, Context>
{
pub fn generate() -> Self {
Self {
main_key: Key::generate(),
context: Context::generate(),
}
}
#[deprecated(note = "use generate() instead")]
pub fn r#gen() -> Self {
Self::generate()
}
}
impl<
Key: ByteArray<CRYPTO_KDF_KEYBYTES> + Zeroize + ZeroizeOnDrop,
Context: ByteArray<CRYPTO_KDF_CONTEXTBYTES> + Zeroize,
> Kdf<Key, Context>
{
pub fn derive_subkey<const LENGTH: usize, Subkey: NewByteArray<LENGTH>>(
&self,
subkey_id: u64,
) -> Result<Subkey, Error> {
validate_subkey_length(LENGTH)?;
let mut subkey = Subkey::new_byte_array();
crypto_kdf_derive_from_key(
subkey.as_mut_array(),
subkey_id,
self.context.as_array(),
self.main_key.as_array(),
)?;
Ok(subkey)
}
pub fn derive_subkey_to_vec(&self, subkey_id: u64, length: usize) -> Result<Vec<u8>, Error> {
validate_subkey_length(length)?;
let mut subkey = vec![0u8; length];
crypto_kdf_derive_from_key(
&mut subkey,
subkey_id,
self.context.as_array(),
self.main_key.as_array(),
)?;
Ok(subkey)
}
pub fn from_parts(main_key: Key, context: Context) -> Self {
Self { main_key, context }
}
pub fn into_parts(self) -> (Key, Context) {
(self.main_key, self.context)
}
}
impl Kdf<Key, Context> {
pub fn generate_with_defaults() -> Self {
Self {
main_key: Key::generate(),
context: Context::generate(),
}
}
#[deprecated(note = "use generate_with_defaults() instead")]
pub fn gen_with_defaults() -> Self {
Self::generate_with_defaults()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_kdf() {
let key = StackKdf::generate();
let short_subkey: StackByteArray<16> = key.derive_subkey(0).expect("derive failed");
let long_subkey = key.derive_subkey_to_vec(0, 64).expect("derive failed");
assert_eq!(short_subkey.len(), 16);
assert_eq!(long_subkey.len(), 64);
assert!(format!("{key:?}").contains("[REDACTED]"));
assert!(matches!(
key.derive_subkey_to_vec(0, usize::MAX),
Err(Error::InvalidLength {
context: crate::ErrorContext::Subkey,
actual: usize::MAX,
..
})
));
let invalid_fixed: Result<StackByteArray<15>, Error> = key.derive_subkey(0);
assert!(invalid_fixed.is_err());
}
}