#[cfg(any(feature = "x25519", feature = "nistp"))]
use crate::dhkex::DhKeyExchange;
use crate::{
Deserializable, Serializable,
aead::{Aead, AeadCtx, AeadCtxR, AeadCtxS, AeadKey, AeadNonce},
kdf::Kdf as KdfTrait,
kem::Kem as KemTrait,
op_mode::{OpModeR, OpModeS, PskBundle},
setup::ExporterSecret,
};
use aead::inout::InOutBuf;
use hybrid_array::Array;
use rand::{Rng, RngExt};
use rand_core::{CryptoRng, TryCryptoRng, TryRng};
pub(crate) fn gen_rand_buf() -> [u8; 32] {
let mut csprng = rand::rng();
let mut buf = [0u8; 32];
csprng.fill_bytes(&mut buf);
buf
}
#[cfg(any(feature = "x25519", feature = "nistp"))]
pub(crate) fn dhkex_gen_keypair_with_rng<Kex: DhKeyExchange>(
csprng: &mut impl CryptoRng,
) -> (Kex::PrivateKey, Kex::PublicKey) {
let mut ikm: Array<u8, <Kex::PrivateKey as Serializable>::OutputSize> = Array::default();
csprng.fill_bytes(&mut ikm);
Kex::derive_keypair::<crate::kdf::HkdfSha512>(b"31337", &ikm)
}
pub(crate) fn gen_ctx_simple_pair<A, Kdf, Kem>() -> (AeadCtxS<A, Kdf, Kem>, AeadCtxR<A, Kdf, Kem>)
where
A: Aead,
Kdf: KdfTrait,
Kem: KemTrait,
{
let mut csprng = rand::rng();
let key = {
let mut buf = AeadKey::<A>::default();
csprng.fill_bytes(buf.0.as_mut_slice());
buf
};
let base_nonce = {
let mut buf = AeadNonce::<A>::default();
csprng.fill_bytes(buf.0.as_mut_slice());
buf
};
let exporter_secret = {
let mut buf = ExporterSecret::<Kdf>::default();
csprng.fill_bytes(buf.0.as_mut_slice());
buf
};
let ctx1 = AeadCtx::new(&key, base_nonce.clone(), exporter_secret.clone());
let ctx2 = AeadCtx::new(&key, base_nonce, exporter_secret);
(ctx1.into(), ctx2.into())
}
#[derive(Clone, Copy)]
pub(crate) enum OpModeKind {
Base,
Auth,
Psk,
AuthPsk,
}
pub(crate) fn new_op_mode_pair<'a, Kem: KemTrait>(
kind: OpModeKind,
psk: &'a [u8],
psk_id: &'a [u8],
) -> (OpModeS<'a, Kem>, OpModeR<'a, Kem>) {
let (sk_sender, pk_sender) = Kem::gen_keypair_with_rng(&mut rand::rng());
let psk_bundle = PskBundle::new(psk, psk_id).unwrap();
match kind {
OpModeKind::Base => {
let sender_mode = OpModeS::Base;
let receiver_mode = OpModeR::Base;
(sender_mode, receiver_mode)
}
OpModeKind::Psk => {
let sender_mode = OpModeS::Psk(psk_bundle);
let receiver_mode = OpModeR::Psk(psk_bundle);
(sender_mode, receiver_mode)
}
OpModeKind::Auth => {
let sender_mode = OpModeS::Auth((sk_sender, pk_sender.clone()));
let receiver_mode = OpModeR::Auth(pk_sender);
(sender_mode, receiver_mode)
}
OpModeKind::AuthPsk => {
let sender_mode = OpModeS::AuthPsk((sk_sender, pk_sender.clone()), psk_bundle);
let receiver_mode = OpModeR::AuthPsk(pk_sender, psk_bundle);
(sender_mode, receiver_mode)
}
}
}
pub(crate) fn aead_ctx_eq<A: Aead, Kdf: KdfTrait, Kem: KemTrait>(
sender: &mut AeadCtxS<A, Kdf, Kem>,
receiver: &mut AeadCtxR<A, Kdf, Kem>,
) -> bool {
let mut csprng = rand::rng();
let msg_len = csprng.random::<u8>() as usize;
let msg_backing_arr = {
let mut buf = [0u8; 255];
csprng.fill_bytes(&mut buf);
buf
};
let msg = &msg_backing_arr[..msg_len];
let aad_len = csprng.random::<u8>() as usize;
let aad_buf = {
let mut buf = [0u8; 255];
csprng.fill_bytes(&mut buf);
buf
};
let aad = &aad_buf[..aad_len];
for i in 0..1000 {
let mut tmp_backing_arr = msg_backing_arr;
let buf = &mut tmp_backing_arr[..msg_len];
let tag = sender
.seal_inout_detached(InOutBuf::new(msg, buf).unwrap(), aad)
.unwrap_or_else(|_| panic!("seal() #{} failed", i));
if receiver
.open_inout_detached(InOutBuf::from(&mut *buf), aad, &tag)
.is_err()
{
return false;
}
if &msg_backing_arr[..msg_len] != buf {
return false;
}
}
true
}
impl Serializable for core::convert::Infallible {
type OutputSize = hybrid_array::typenum::U0;
fn write_exact(&self, _: &mut [u8]) {
unimplemented!()
}
}
impl Deserializable for core::convert::Infallible {
fn from_bytes(_: &[u8]) -> Result<Self, crate::HpkeError> {
unimplemented!()
}
}
pub(crate) struct FakeCsprng<'a> {
randomness: &'a [u8],
}
impl<'a> FakeCsprng<'a> {
pub(crate) fn new(randomness: &'a [u8]) -> Self {
Self { randomness }
}
}
impl<'a> TryRng for FakeCsprng<'a> {
type Error = core::convert::Infallible;
fn try_next_u32(&mut self) -> Result<u32, Self::Error> {
rand_core::utils::next_word_via_fill(self)
}
fn try_next_u64(&mut self) -> Result<u64, Self::Error> {
rand_core::utils::next_word_via_fill(self)
}
fn try_fill_bytes(&mut self, dest: &mut [u8]) -> Result<(), Self::Error> {
if dest.len() > self.randomness.len() {
unreachable!("ran out of randomness")
} else {
let (taken, rest) = self.randomness.split_at(dest.len());
dest.copy_from_slice(taken);
self.randomness = rest;
Ok(())
}
}
}
impl<'a> TryCryptoRng for FakeCsprng<'a> {}