responder side of the handshake and full e2e test

This commit is contained in:
Jędrzej Stuczyński
2026-02-17 11:24:59 +00:00
parent 4fcb8ed202
commit 26056909b7
10 changed files with 611 additions and 975 deletions
+2
View File
@@ -7,6 +7,8 @@ use std::collections::HashMap;
use std::fmt::Display;
use strum_macros::{Display, EnumIter, EnumString};
pub use strum::IntoEnumIterator;
pub mod error;
pub const DEFAULT_HASH_LEN: usize = 32;
+7 -1
View File
@@ -8,7 +8,7 @@ use crate::frame::KKTFrame;
use crate::keys::EncapsulationKey;
use crate::masked_byte::{MASKED_BYTE_LEN, MaskedByte};
use libcrux_psq::handshake::types::{DHKeyPair, DHPrivateKey, DHPublicKey};
use nym_kkt_ciphersuite::x25519;
use nym_kkt_ciphersuite::{KEM, x25519};
pub struct KKTRequest {
/// The plaintext part of the request
@@ -185,6 +185,12 @@ pub struct ProcessedKKTRequest {
/// The obtained encapsulation key of the remote
pub remote_encapsulation_key: Option<EncapsulationKey>,
/// The KEM key requested in the original request
pub requested_kem: KEM,
/// The unmasked byte representing the outer protocol version sent by the initiator
pub outer_protocol_version: u8,
}
pub struct KKTResponse {
+2
View File
@@ -127,6 +127,8 @@ impl<'a> KKTResponder<'a> {
Ok(ProcessedKKTRequest {
response: KKTResponse { encrypted_frame },
remote_encapsulation_key,
requested_kem: remote_context.ciphersuite().kem(),
outer_protocol_version: processed_req.outer_protocol_version,
})
}
}
+4 -4
View File
@@ -100,8 +100,8 @@ pub enum LpError {
#[error("incompatible LP packet version. got: {got}, lowest supported: {lowest_supported}")]
IncompatibleLegacyPacketVersion { got: u8, lowest_supported: u8 },
#[error("attempted to create an LP responder without providing a valid KEM key for {kem} ")]
ResponderWithMissingKEMKey { kem: KEM },
#[error("attempted to create an LP responder without providing a valid KEM keys")]
ResponderWithMissingKEMKeys,
#[error(
"there are no known digests for remote's KEM key with {kem} KEM and {hash_function} hash function"
@@ -135,8 +135,8 @@ pub enum LpError {
#[error("failed to run the PSQ session: {inner:?}")]
PSQSessionFailure { inner: SessionError },
#[error("can't proceed without remote peer information")]
MissingRemotePeerInformation,
#[error("the initiator authenticator is not available after ingesting PSQ msg1")]
MissingInitiatorAuthenticator,
}
impl LpError {
+218 -86
View File
@@ -8,7 +8,7 @@ use crate::peer::{LpLocalPeer, LpRemotePeer};
use crate::psk::psq_initiator_create_message;
use crate::psq::helpers::{LpTransportHandshakeExt, current_timestamp, kem_to_ciphersuite};
use crate::psq::{
AAD_INITIATOR_INNER_V1, AAD_INITIATOR_OUTER_V1, IntermediateHandshakeFailure, MinimalSession,
AAD_INITIATOR_INNER_V1, AAD_INITIATOR_OUTER_V1, InitiatorData, MinimalSession,
PSQHandshakeState, SESSION_CONTEXT_V1, initiator,
};
use crate::session::PqSharedSecret;
@@ -28,6 +28,11 @@ use nym_lp_transport::traits::LpTransport;
use rand09::rng;
use tracing::debug;
pub(crate) struct PSQHandshakeStateInitiator<'a, S> {
pub(super) inner_state: PSQHandshakeState<'a, S>,
pub(super) initiator_data: InitiatorData,
}
pub(crate) fn build_psq_principal<R>(
rng: R,
version: u8,
@@ -72,7 +77,7 @@ pub(crate) fn build_psq_ciphersuite<'a>(
.map_err(|inner| LpError::PSQInitiatorBuilderFailure { inner })
}
impl<'a, S> PSQHandshakeState<'a, S>
impl<'a, S> PSQHandshakeStateInitiator<'a, S>
where
S: LpTransport + Unpin,
{
@@ -80,96 +85,53 @@ where
&'b self,
encapsulation_key: &'b EncapsulationKey,
) -> Result<RegistrationInitiator<'b, rand09::rngs::ThreadRng>, LpError> {
let initiator_ciphersuite =
build_psq_ciphersuite(&self.local_peer, self.remote_peer()?, &encapsulation_key)?;
let initiator =
build_psq_principal(rng(), self.protocol_version()?, initiator_ciphersuite)?;
let initiator_ciphersuite = build_psq_ciphersuite(
&self.inner_state.local_peer,
&self.initiator_data.remote_peer,
&encapsulation_key,
)?;
let initiator = build_psq_principal(
rng(),
self.initiator_data.protocol_version,
initiator_ciphersuite,
)?;
Ok(initiator)
}
/// Attempt to send KKT request to begin the handshake
pub(crate) async fn send_kkt_request(&mut self, request: KKTRequest) -> Result<(), LpError> {
async fn send_kkt_request(&mut self, request: KKTRequest) -> Result<(), LpError> {
// TODO: extra header
self.connection
self.inner_state
.connection
.send_serialised_packet(&request.into_bytes())
.await?;
Ok(())
}
/// Attempt to receive a KKT response to the previously sent request
pub(crate) async fn receive_kkt_response(&mut self) -> Result<KKTResponse, LpError> {
let data = self.connection.receive_raw_packet().await?;
async fn receive_kkt_response(&mut self) -> Result<KKTResponse, LpError> {
let data = self.inner_state.connection.receive_raw_packet().await?;
Ok(KKTResponse::from_bytes(data))
}
/// Attempt to prepare and send final PSQ msg3
pub(crate) async fn send_final_psq_message(
&mut self,
session_id: u32,
outer_aead_key: &OuterAeadKey,
noise_protocol: &mut NoiseProtocol,
) -> Result<(), LpError> {
todo!()
// let protocol = self.protocol_version()?;
//
// let noise_msg3 = noise_protocol
// .get_bytes_to_send()
// .ok_or_else(|| LpError::kkt_psq_handshake("failed to generate noise msg3"))??;
//
// let lp_message = HandshakeData::new(noise_msg3).into();
// let lp_packet = self.next_packet(session_id, protocol, lp_message);
// self.connection
// .send_packet(lp_packet, Some(outer_aead_key))
// .await?;
//
// if !noise_protocol.is_handshake_finished() {
// return Err(LpError::kkt_psq_handshake(
// "noise handshake not finished after msg3",
// ));
// }
//
// Ok(())
}
/// Receive final ACK that indicates finalisation of the handshake
pub(crate) async fn receive_final_ack(
&mut self,
outer_aead_key: &OuterAeadKey,
) -> Result<(), LpError> {
match self
.connection
.receive_packet(Some(outer_aead_key))
.await?
.message
{
LpMessage::Ack => Ok(()),
other => Err(LpError::unexpected_handshake_response(
other.typ(),
MessageType::Ack,
)),
}
}
pub async fn complete_as_initiator_inner<R>(
mut self,
rng: &mut R,
) -> Result<MinimalSession, LpError>
pub async fn complete_handshake<R>(mut self, rng: &mut R) -> Result<MinimalSession, LpError>
where
S: LpTransport + Unpin,
R: rand09::CryptoRng,
{
// 1. retrieve the expected kem key hash. if we don't know it,
let dir_hash = self
.remote_peer()?
.expected_kem_key_hash(self.ciphersuite)?;
.initiator_data
.remote_peer
.expected_kem_key_hash(self.inner_state.ciphersuite)?;
// 2. prepare and send KKT request
let (mut initiator, kkt_request) = KKTInitiator::generate_one_way_request(
rng,
self.ciphersuite,
self.remote_peer()?.x25519(),
self.inner_state.ciphersuite,
self.initiator_data.remote_peer.x25519(),
&dir_hash,
self.protocol_version()?,
self.initiator_data.protocol_version,
)?;
debug!("sending KKT request");
self.send_kkt_request(kkt_request).await?;
@@ -180,17 +142,15 @@ where
let response = initiator.process_response(raw_response)?;
// 4. generate and send PSQ request
let protocol = self.protocol_version()?;
let mut conn = self.connection;
let remote_peer = self
.remote_peer
.as_ref()
.ok_or(LpError::MissingRemotePeerInformation)?;
let protocol = self.initiator_data.protocol_version;
let mut conn = self.inner_state.connection;
// build the PSQ initiator
let initiator_ciphersuite =
build_psq_ciphersuite(&self.local_peer, remote_peer, &response.encapsulation_key)?;
let initiator_ciphersuite = build_psq_ciphersuite(
&self.inner_state.local_peer,
&self.initiator_data.remote_peer,
&response.encapsulation_key,
)?;
let mut psq_initiator = build_psq_principal(rng, protocol, initiator_ciphersuite)?;
@@ -218,18 +178,190 @@ where
Ok(MinimalSession {
session,
encapsulation_key: Some(response.encapsulation_key),
init_authenticator: None,
})
}
}
// TODO: missing: receive counter check
pub async fn complete_as_initiator(mut self) -> Result<LpSession, LpError>
where
S: LpTransport + Unpin,
{
todo!()
// match self.complete_as_initiator_inner().await {
// Ok(res) => Ok(res),
// Err(err) => Err(self.try_send_error_packet(err).await),
// }
#[cfg(test)]
mod tests {
use super::*;
use crate::peer::mock_peers;
use crate::psq::responder;
use libcrux_psq::handshake::types::Authenticator;
use libcrux_psq::session::{Session, SessionBinding};
use nym_kkt::responder::KKTResponder;
use nym_kkt_ciphersuite::{Ciphersuite, HashFunction, SignatureScheme};
use nym_test_utils::helpers::{DeterministicRng09Send, u64_seeded_rng_09};
use nym_test_utils::mocks::async_read_write::MockIOStream;
use nym_test_utils::traits::{Leak, Timeboxed};
#[tokio::test]
async fn initiator_test_plain() -> anyhow::Result<()> {
let conn_init = MockIOStream::default();
let conn_resp = conn_init.try_get_remote_handle();
// leak the connections (JUST FOR THE PURPOSE OF THIS TEST!)
// so they'd get 'static lifetime
let conn_init = conn_init.leak();
let conn_resp = conn_resp.leak();
let (init, resp) = mock_peers();
let init_remote = init.as_remote();
let resp_remote = resp.as_remote();
let kem = KEM::MlKem768;
let ciphersuite = Ciphersuite::default().with_kem(kem);
let initiator_data = InitiatorData::new(1, resp_remote);
let handshake_init =
PSQHandshakeState::new(conn_init, ciphersuite, init).as_initiator(initiator_data);
let mut init_rng = DeterministicRng09Send::new(u64_seeded_rng_09(1));
let init_fut = tokio::spawn(async move {
handshake_init
.complete_handshake(&mut init_rng)
.timeboxed()
.await
});
// responder:
let supported_sigs = [SignatureScheme::Ed25519];
let supported_hash = [
HashFunction::Blake3,
HashFunction::Shake256,
HashFunction::Shake128,
HashFunction::SHA256,
];
let resp_keys = resp.kem_keypairs.as_ref().unwrap();
let responder_x25519_keypair = resp.x25519();
let kkt_responder = KKTResponder::new(
&responder_x25519_keypair,
&resp_keys,
&supported_hash,
&supported_sigs,
&[1],
)?;
// 1. read KKT request
let raw_kkt_req = conn_resp.receive_raw_packet().timeboxed().await??;
let req = KKTRequest::try_from_bytes(&raw_kkt_req)?;
// 2. process
let processed_req = kkt_responder.process_request(req)?;
conn_resp
.send_serialised_packet(&processed_req.response.into_bytes())
.timeboxed()
.await??;
// 3. read PSQ req
let responder_ciphersuite = responder::build_psq_ciphersuite(&resp, kem)?;
let mut responder =
responder::build_psq_principal(rand09::rng(), 1, responder_ciphersuite)?;
let raw_psq_req = conn_resp.receive_raw_packet().timeboxed().await??;
responder.read_message(&raw_psq_req, &mut []).unwrap();
// Get the authenticator out here, so we can deserialize the session later.
let Some(initiator_authenticator) = responder.initiator_authenticator() else {
panic!("No initiator authenticator found")
};
// 4 send PSQ response
let mut buf = [0u8; 2048];
let n = responder.write_message(&[], &mut buf).unwrap();
conn_resp
.send_serialised_packet(&buf[..n])
.timeboxed()
.await??;
assert!(responder.is_handshake_finished());
let session_init = init_fut.await???;
let i_transport = session_init.session;
let encapsulation_key = session_init.encapsulation_key.unwrap();
let r_transport = responder.into_session().unwrap();
// test serialization, deserialization
let mut msg_channel = vec![0u8; 2048];
let mut payload_buf_responder = vec![0u8; 4096];
let mut payload_buf_initiator = vec![0u8; 4096];
let mut session_storage = vec![0u8; 4096];
i_transport
.serialize(
&mut session_storage,
SessionBinding {
initiator_authenticator: &Authenticator::Dh(init_remote.x25519_public),
responder_ecdh_pk: &responder_x25519_keypair.pk,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
let mut i_transport = Session::deserialize(
&session_storage,
SessionBinding {
initiator_authenticator: &Authenticator::Dh(init_remote.x25519_public),
responder_ecdh_pk: &responder_x25519_keypair.pk,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
r_transport
.serialize(
&mut session_storage,
SessionBinding {
initiator_authenticator: &initiator_authenticator,
responder_ecdh_pk: &responder_x25519_keypair.pk,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
let mut r_transport = Session::deserialize(
&session_storage,
SessionBinding {
initiator_authenticator: &initiator_authenticator,
responder_ecdh_pk: &responder_x25519_keypair.pk,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
let mut channel_i = i_transport.transport_channel().unwrap();
let mut channel_r = r_transport.transport_channel().unwrap();
assert_eq!(channel_i.identifier(), channel_r.identifier());
let app_data_i = b"Derived session hey".as_slice();
let app_data_r = b"Derived session ho".as_slice();
let len_i = channel_i
.write_message(app_data_i, &mut msg_channel)
.unwrap();
let (len_r_deserialized, len_r_payload) = channel_r
.read_message(&msg_channel, &mut payload_buf_responder)
.unwrap();
// We read the same amount of data.
assert_eq!(len_r_deserialized, len_i);
assert_eq!(len_r_payload, app_data_i.len());
assert_eq!(&payload_buf_responder[0..len_r_payload], app_data_i);
let len_r = channel_r
.write_message(app_data_r, &mut msg_channel)
.unwrap();
let (len_i_deserialized, len_i_payload) = channel_i
.read_message(&msg_channel, &mut payload_buf_initiator)
.unwrap();
assert_eq!(len_r, len_i_deserialized);
assert_eq!(app_data_r.len(), len_i_payload);
assert_eq!(&payload_buf_initiator[0..len_i_payload], app_data_r);
Ok(())
}
}
+99 -439
View File
@@ -2,12 +2,16 @@
// SPDX-License-Identifier: Apache-2.0
use crate::codec::OuterAeadKey;
use crate::packet::version;
use crate::peer::{LpLocalPeer, LpRemotePeer};
use crate::psq::helpers::LpTransportHandshakeExt;
use crate::psq::initiator::PSQHandshakeStateInitiator;
use crate::psq::responder::PSQHandshakeStateResponder;
use crate::{LpError, LpMessage};
use libcrux_psq::handshake::types::Authenticator;
use libcrux_psq::session::Session;
use nym_kkt::keys::EncapsulationKey;
use nym_kkt_ciphersuite::Ciphersuite;
use nym_kkt_ciphersuite::{Ciphersuite, HashFunction, IntoEnumIterator, SignatureScheme};
use nym_lp_transport::traits::LpTransport;
mod helpers;
@@ -22,32 +26,7 @@ pub(crate) const SESSION_CONTEXT_V1: &[u8] = b"NYM-PQ-SESSION-CONTEXT-V1";
pub struct MinimalSession {
session: Session,
encapsulation_key: Option<EncapsulationKey>,
}
#[deprecated]
pub(crate) struct IntermediateHandshakeFailure {
/// Session id established during exchange if we managed to derive it
session_id: Option<u32>,
/// Protocol version established during the exchange
protocol_version: Option<u8>,
/// Outer aead key established during exchange if we managed to derive it
outer_aead_key: Option<OuterAeadKey>,
/// The error source
source: LpError,
}
impl IntermediateHandshakeFailure {
fn plain(source: LpError) -> IntermediateHandshakeFailure {
IntermediateHandshakeFailure {
session_id: None,
protocol_version: None,
outer_aead_key: None,
source,
}
}
init_authenticator: Option<Authenticator>,
}
pub struct PSQHandshakeState<'a, S> {
@@ -65,10 +44,48 @@ pub struct PSQHandshakeState<'a, S> {
/// Representation of a local Lewes Protocol peer
/// encapsulating all the known information and keys.
local_peer: LpLocalPeer,
}
#[derive(Debug)]
pub struct InitiatorData {
/// Protocol version used for the exchange known implicitly through the directory
pub protocol_version: u8,
/// Representation of a remote Lewes Protocol peer
/// encapsulating all the known information and keys.
remote_peer: Option<LpRemotePeer>,
pub remote_peer: LpRemotePeer,
}
impl InitiatorData {
pub fn new(protocol_version: u8, remote_peer: LpRemotePeer) -> Self {
InitiatorData {
protocol_version,
remote_peer,
}
}
}
#[derive(Debug, Clone)]
pub struct ResponderData {
/// List of supported Hash Functions by this Responder
pub supported_hash_functions: Vec<HashFunction>,
/// List of supported Signature Schemes by this Responder
pub supported_signature_schemes: Vec<SignatureScheme>,
/// List of supported outer (LP) protocol version by this Responder
pub supported_outer_protocol_versions: Vec<u8>,
}
impl Default for ResponderData {
fn default() -> Self {
// by default all schemes are supported
ResponderData {
supported_hash_functions: HashFunction::iter().collect(),
supported_signature_schemes: SignatureScheme::iter().collect(),
supported_outer_protocol_versions: vec![version::CURRENT],
}
}
}
impl<'a, S> PSQHandshakeState<'a, S>
@@ -81,100 +98,22 @@ where
protocol_version: None,
ciphersuite,
local_peer,
remote_peer: None,
}
}
#[must_use]
pub fn with_protocol_version(mut self, protocol_version: u8) -> Self {
self.protocol_version = Some(protocol_version);
self
pub fn as_initiator(self, initiator_data: InitiatorData) -> PSQHandshakeStateInitiator<'a, S> {
PSQHandshakeStateInitiator {
initiator_data,
inner_state: self,
}
}
#[must_use]
pub fn with_remote_peer(mut self, remote_peer: LpRemotePeer) -> Self {
self.remote_peer = Some(remote_peer);
self
pub fn as_responder(self, responder_data: ResponderData) -> PSQHandshakeStateResponder<'a, S> {
PSQHandshakeStateResponder {
responder_data,
inner_state: self,
}
}
fn protocol_version(&self) -> Result<u8, LpError> {
self.protocol_version
.ok_or_else(|| LpError::kkt_psq_handshake("unknown protocol version"))
}
fn remote_peer(&self) -> Result<&LpRemotePeer, LpError> {
self.remote_peer
.as_ref()
.ok_or(LpError::MissingRemotePeerInformation)
}
//
// pub fn next_packet(
// &mut self,
// session_id: u32,
// protocol_version: u8,
// message: LpMessage,
// ) -> LpPacket {
// let counter = self.next_counter();
// let header = LpHeader::new(session_id, counter, protocol_version);
// LpPacket::new(header, message)
// }
//
// pub(crate) async fn try_send_error_packet(
// &mut self,
// err: IntermediateHandshakeFailure,
// ) -> LpError {
// // if session_id is not known, we can't send the packet back (with the current design)
// let (Some(session_id), Some(protocol)) = (err.session_id, err.protocol_version) else {
// return err.source;
// };
// if let Err(err) = self
// .send_error_packet(
// session_id,
// protocol,
// err.source.to_string(),
// err.outer_aead_key.as_ref(),
// )
// .await
// {
// debug!("failed to send back error response: {err}")
// }
// err.source
// }
//
// /// Attempt to send an error packet
// pub(crate) async fn send_error_packet(
// &mut self,
// session_id: u32,
// protocol_version: u8,
// msg: impl Into<String>,
// outer_aead_key: Option<&OuterAeadKey>,
// ) -> Result<(), LpError> {
// let packet = self.next_packet(
// session_id,
// protocol_version,
// LpMessage::Error(ErrorPacketData::new(msg)),
// );
// self.connection.send_packet(packet, outer_aead_key).await?;
// Ok(())
// }
//
// /// Attempt to receive a packet from connection, explicitly checking for an error response
// /// and returning corresponding message if received
// pub(crate) async fn receive_non_error(
// &mut self,
// outer_aead_key: Option<&OuterAeadKey>,
// ) -> Result<LpPacket, LpError> {
// let packet = self.connection.receive_packet(outer_aead_key).await?;
//
// match &packet.message {
// LpMessage::Error(error_packet) => Err(LpError::kkt_psq_handshake(format!(
// "remote error: {}",
// error_packet.message
// ))),
// _ => Ok(packet),
// }
// }
}
#[cfg(test)]
@@ -182,26 +121,18 @@ mod tests {
use super::*;
use crate::peer::mock_peers;
use crate::psq::helpers::LpTransportHandshakeExt;
use crate::psq::responder::DEFAULT_TIMESTAMP_TOLERANCE;
use libcrux_psq::handshake::types::{Authenticator, PQEncapsulationKey};
use libcrux_psq::handshake::types::Authenticator;
use libcrux_psq::session::{Session, SessionBinding};
use libcrux_psq::{Channel, IntoSession};
use mock_instant::thread_local::MockClock;
use nym_kkt::initiator::KKTInitiator;
use nym_kkt::key_utils::{
generate_keypair_mceliece, generate_keypair_mlkem, generate_keypair_x25519,
hash_encapsulation_key,
};
use nym_kkt::keys::EncapsulationKey;
use nym_kkt::message::KKTRequest;
use nym_kkt::message::{KKTRequest, KKTResponse};
use nym_kkt::responder::KKTResponder;
use nym_kkt_ciphersuite::{HashFunction, HashLength, KEM, SignatureScheme};
use nym_kkt_ciphersuite::{HashFunction, KEM, SignatureScheme};
use nym_test_utils::helpers::{
DeterministicRng09Send, deterministic_rng_09, u64_seeded_rng_09,
};
use nym_test_utils::mocks::async_read_write::MockIOStream;
use nym_test_utils::traits::{Leak, Timeboxed, TimeboxedSpawnable};
use std::time::Duration;
use tokio::join;
#[allow(dead_code)]
@@ -215,151 +146,50 @@ mod tests {
#[tokio::test]
async fn e2e_psq_handshake() -> anyhow::Result<()> {
todo!()
// let conn_init = MockIOStream::default();
// let conn_resp = conn_init.try_get_remote_handle();
//
// // leak the connections (JUST FOR THE PURPOSE OF THIS TEST!)
// // so they'd get 'static lifetime
// let conn_init = conn_init.leak();
// let conn_resp = conn_resp.leak();
//
// let ciphersuite = Ciphersuite::new(
// KEM::X25519,
// HashFunction::Blake3,
// SignatureScheme::Ed25519,
// HashLength::Default,
// );
//
// let (init, resp) = mock_peers();
// let resp_remote = resp.as_remote();
//
// let handshake_init = PSQHandshakeState::new(conn_init, ciphersuite, init)
// .with_protocol_version(1)
// .with_remote_peer(resp_remote);
// let handshake_resp = PSQHandshakeState::new(conn_resp, ciphersuite, resp);
//
// let resp_fut = handshake_resp.complete_as_responder().spawn_timeboxed();
// let init_fut = handshake_init.complete_as_initiator().spawn_timeboxed();
//
// let (session_init, session_resp) = join!(init_fut, resp_fut);
//
// let session_init = session_init???;
// let session_resp = session_resp???;
//
// assert_eq!(session_init.id(), session_resp.id());
// assert_eq!(
// session_init.outer_aead_key().as_bytes(),
// session_resp.outer_aead_key().as_bytes()
// );
// assert_eq!(
// session_init.pq_shared_secret().as_bytes(),
// session_resp.pq_shared_secret().as_bytes()
// );
//
// Ok(())
}
let conn_init = MockIOStream::default();
let conn_resp = conn_init.try_get_remote_handle();
#[tokio::test]
async fn preparing_client_hello_initiator() -> anyhow::Result<()> {
todo!()
// let mut conn_init = MockIOStream::default();
// let mut conn_resp = conn_init.try_get_remote_handle();
//
// let ciphersuite = Ciphersuite::new(
// KEM::X25519,
// HashFunction::Blake3,
// SignatureScheme::Ed25519,
// HashLength::Default,
// );
// let (init, resp) = mock_peers();
// let resp_remote = resp.as_remote();
//
// // as initiator
// let mut handshake_init = PSQHandshakeState::new(&mut conn_init, ciphersuite, init)
// .with_protocol_version(1)
// .with_remote_peer(resp_remote);
//
// // you can generate and send (valid) client hello as initiator
// let client_hello = handshake_init.send_client_hello().await?;
// let LpMessage::ClientHello(received_client_hello) =
// conn_resp.receive_packet(None).await?.message
// else {
// panic!("wrong message type");
// };
// assert_eq!(client_hello, received_client_hello);
// Ok(())
}
// leak the connections (JUST FOR THE PURPOSE OF THIS TEST!)
// so they'd get 'static lifetime
let conn_init = conn_init.leak();
let conn_resp = conn_resp.leak();
// essentially make sure you can't accidentally trigger the handshake as the responder
#[tokio::test]
async fn preparing_client_hello_responder() -> anyhow::Result<()> {
todo!()
// let conn_init = MockIOStream::default();
// let mut conn_resp = conn_init.try_get_remote_handle();
//
// let ciphersuite = Ciphersuite::new(
// KEM::X25519,
// HashFunction::Blake3,
// SignatureScheme::Ed25519,
// HashLength::Default,
// );
// let (_, resp) = mock_peers();
//
// // as initiator
// let mut handshake_resp = PSQHandshakeState::new(&mut conn_resp, ciphersuite, resp);
//
// // you can generate and send (valid) client hello as initiator
// let sending_res = handshake_resp.send_client_hello().await;
// assert!(sending_res.is_err());
// Ok(())
}
let kem = KEM::MlKem768;
let ciphersuite = Ciphersuite::default().with_kem(kem);
#[tokio::test]
async fn test_receive_client_hello_timestamp_too_skewed() -> anyhow::Result<()> {
todo!()
// let current_time = Duration::from_secs(10000);
// MockClock::set_system_time(current_time);
//
// let too_old = current_time - DEFAULT_TIMESTAMP_TOLERANCE - Duration::from_secs(1);
// let too_recent = current_time + DEFAULT_TIMESTAMP_TOLERANCE + Duration::from_secs(1);
//
// let ciphersuite = Ciphersuite::new(
// KEM::X25519,
// HashFunction::Blake3,
// SignatureScheme::Ed25519,
// HashLength::Default,
// );
//
// // TOO OLD
// let mut conn_init = MockIOStream::default();
// let mut conn_resp = conn_init.try_get_remote_handle();
// let (init, resp) = mock_peers();
//
// let mut handshake_resp = PSQHandshakeState::new(&mut conn_resp, ciphersuite, resp);
// let client_hello_too_old = init.build_client_hello_data(too_old.as_secs());
//
// conn_init
// .send_packet(client_hello_too_old.into_lp_packet(1), None)
// .await?;
// let err = handshake_resp.receive_client_hello().await.unwrap_err();
// assert!(err.to_string().contains("too old"));
//
// // TOO RECENT
// let mut conn_init = MockIOStream::default();
// let mut conn_resp = conn_init.try_get_remote_handle();
// let (init, resp) = mock_peers();
//
// let mut handshake_resp = PSQHandshakeState::new(&mut conn_resp, ciphersuite, resp);
// let client_hello_too_recent = init.build_client_hello_data(too_recent.as_secs());
//
// conn_init
// .send_packet(client_hello_too_recent.into_lp_packet(1), None)
// .await?;
// let err = handshake_resp.receive_client_hello().await.unwrap_err();
//
// assert!(err.to_string().contains("too future"));
// Ok(())
let (init, resp) = mock_peers();
let resp_remote = resp.as_remote();
let handshake_init = PSQHandshakeState::new(conn_init, ciphersuite, init)
.as_initiator(InitiatorData::new(1, resp_remote));
let handshake_resp = PSQHandshakeState::new(conn_resp, ciphersuite, resp)
.as_responder(ResponderData::default());
let init_rng = DeterministicRng09Send::new(u64_seeded_rng_09(1));
let resp_rng = DeterministicRng09Send::new(u64_seeded_rng_09(2));
// similarly leak the rngs to get the static lifetimes
let init_rng = init_rng.leak();
let resp_rng = resp_rng.leak();
let init_fut = handshake_init
.complete_handshake(init_rng)
.spawn_timeboxed();
let resp_fut = handshake_resp
.complete_handshake(resp_rng)
.spawn_timeboxed();
let (session_init, session_resp) = join!(init_fut, resp_fut);
let session_init = session_init???;
let session_resp = session_resp???;
assert_eq!(
session_init.session.identifier(),
session_resp.session.identifier()
);
Ok(())
}
// plain test without any wrappers
@@ -532,174 +362,4 @@ mod tests {
assert_eq!(app_data_r.len(), len_i_payload);
assert_eq!(&payload_buf_initiator[0..len_i_payload], app_data_r);
}
#[tokio::test]
async fn initiator_test_plain() -> anyhow::Result<()> {
let conn_init = MockIOStream::default();
let conn_resp = conn_init.try_get_remote_handle();
// leak the connections (JUST FOR THE PURPOSE OF THIS TEST!)
// so they'd get 'static lifetime
let conn_init = conn_init.leak();
let conn_resp = conn_resp.leak();
let (init, resp) = mock_peers();
let init_remote = init.as_remote();
let resp_remote = resp.as_remote();
let kem = KEM::MlKem768;
let ciphersuite = Ciphersuite::default().with_kem(kem);
let handshake_init = PSQHandshakeState::new(conn_init, ciphersuite, init)
.with_protocol_version(1)
.with_remote_peer(resp_remote);
let mut init_rng = DeterministicRng09Send::new(u64_seeded_rng_09(1));
let init_fut = tokio::spawn(async move {
handshake_init
.complete_as_initiator_inner(&mut init_rng)
.timeboxed()
.await
});
// responder:
let supported_sigs = [SignatureScheme::Ed25519];
let supported_hash = [
HashFunction::Blake3,
HashFunction::Shake256,
HashFunction::Shake128,
HashFunction::SHA256,
];
let resp_keys = resp.kem_keypairs.as_ref().unwrap();
let responder_x25519_keypair = resp.x25519();
let kkt_responder = KKTResponder::new(
&responder_x25519_keypair,
&resp_keys,
&supported_hash,
&supported_sigs,
&[1],
)?;
// 1. read KKT request
let raw_kkt_req = conn_resp.receive_raw_packet().timeboxed().await??;
let req = KKTRequest::try_from_bytes(&raw_kkt_req)?;
// 2. process
let processed_req = kkt_responder.process_request(req)?;
conn_resp
.send_serialised_packet(&processed_req.response.into_bytes())
.timeboxed()
.await??;
// 3. read PSQ req
let responder_ciphersuite = responder::build_psq_ciphersuite(&resp, kem)?;
let mut responder =
responder::build_psq_principal(rand09::rng(), 1, responder_ciphersuite)?;
let raw_psq_req = conn_resp.receive_raw_packet().timeboxed().await??;
let mut buf = [0u8; 2048];
responder.read_message(&raw_psq_req, &mut buf).unwrap();
// Get the authenticator out here, so we can deserialize the session later.
let Some(initiator_authenticator) = responder.initiator_authenticator() else {
panic!("No initiator authenticator found")
};
// 4 send PSQ response
let mut buf = [0u8; 2048];
let n = responder.write_message(&[], &mut buf).unwrap();
conn_resp
.send_serialised_packet(&buf[..n])
.timeboxed()
.await??;
assert!(responder.is_handshake_finished());
let session_init = init_fut.await???;
let i_transport = session_init.session;
let encapsulation_key = session_init.encapsulation_key.unwrap();
let r_transport = responder.into_session().unwrap();
// test serialization, deserialization
let mut msg_channel = vec![0u8; 2048];
let mut payload_buf_responder = vec![0u8; 4096];
let mut payload_buf_initiator = vec![0u8; 4096];
let mut session_storage = vec![0u8; 4096];
i_transport
.serialize(
&mut session_storage,
SessionBinding {
initiator_authenticator: &Authenticator::Dh(init_remote.x25519_public),
responder_ecdh_pk: &responder_x25519_keypair.pk,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
let mut i_transport = Session::deserialize(
&session_storage,
SessionBinding {
initiator_authenticator: &Authenticator::Dh(init_remote.x25519_public),
responder_ecdh_pk: &responder_x25519_keypair.pk,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
r_transport
.serialize(
&mut session_storage,
SessionBinding {
initiator_authenticator: &initiator_authenticator,
responder_ecdh_pk: &responder_x25519_keypair.pk,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
let mut r_transport = Session::deserialize(
&session_storage,
SessionBinding {
initiator_authenticator: &initiator_authenticator,
responder_ecdh_pk: &responder_x25519_keypair.pk,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
let mut channel_i = i_transport.transport_channel().unwrap();
let mut channel_r = r_transport.transport_channel().unwrap();
assert_eq!(channel_i.identifier(), channel_r.identifier());
let app_data_i = b"Derived session hey".as_slice();
let app_data_r = b"Derived session ho".as_slice();
let len_i = channel_i
.write_message(app_data_i, &mut msg_channel)
.unwrap();
let (len_r_deserialized, len_r_payload) = channel_r
.read_message(&msg_channel, &mut payload_buf_responder)
.unwrap();
// We read the same amount of data.
assert_eq!(len_r_deserialized, len_i);
assert_eq!(len_r_payload, app_data_i.len());
assert_eq!(&payload_buf_responder[0..len_r_payload], app_data_i);
let len_r = channel_r
.write_message(app_data_r, &mut msg_channel)
.unwrap();
let (len_i_deserialized, len_i_payload) = channel_i
.read_message(&msg_channel, &mut payload_buf_initiator)
.unwrap();
assert_eq!(len_r, len_i_deserialized);
assert_eq!(app_data_r.len(), len_i_payload);
assert_eq!(&payload_buf_initiator[0..len_i_payload], app_data_r);
Ok(())
}
}
+270 -438
View File
@@ -2,27 +2,24 @@
// SPDX-License-Identifier: Apache-2.0
use crate::codec::OuterAeadKey;
use crate::message::{HandshakeData, KKTResponseData, MessageType};
use crate::noise_protocol::NoiseProtocol;
use crate::peer::{LpLocalPeer, LpRemotePeer};
use crate::psk::psq_responder_process_message;
use crate::psq::helpers::{LpTransportHandshakeExt, current_timestamp, kem_to_ciphersuite};
use crate::psq::helpers::kem_to_ciphersuite;
use crate::psq::{
AAD_RESPONDER_V1, IntermediateHandshakeFailure, PSQHandshakeState, SESSION_CONTEXT_V1,
AAD_RESPONDER_V1, MinimalSession, PSQHandshakeState, ResponderData, SESSION_CONTEXT_V1,
};
use crate::session::PqSharedSecret;
use crate::{ClientHelloData, LpError, LpMessage, LpSession};
use crate::{ClientHelloData, LpError, LpSession};
use libcrux_psq::handshake::Responder;
use libcrux_psq::handshake::builders::{
CiphersuiteBuilder, PrincipalBuilder, ResponderCiphersuite,
};
use libcrux_psq::handshake::ciphersuites::CiphersuiteName;
use libcrux_psq::{Channel, IntoSession};
use nym_kkt::context::KKTContext;
use nym_kkt::keys::KEMKeys;
use nym_kkt::message::{KKTRequest, KKTResponse, ProcessedKKTRequest};
use nym_kkt::responder::KKTResponder;
use nym_kkt_ciphersuite::KEM;
use nym_lp_transport::traits::LpTransport;
use rand09::rng;
use std::time::Duration;
use tracing::debug;
pub(crate) fn build_psq_principal<R>(
@@ -51,7 +48,7 @@ pub(crate) fn build_psq_ciphersuite(
kem: KEM,
) -> Result<ResponderCiphersuite, LpError> {
let Some(kem_keys) = peer.kem_keypairs.as_ref() else {
return Err(LpError::ResponderWithMissingKEMKey { kem });
return Err(LpError::ResponderWithMissingKEMKeys);
};
let psq_ciphersuite = kem_to_ciphersuite(kem);
@@ -69,461 +66,296 @@ pub(crate) fn build_psq_ciphersuite(
.map_err(|inner| LpError::PSQResponderBuilderFailure { inner })
}
pub const DEFAULT_TIMESTAMP_TOLERANCE: Duration = Duration::from_secs(30);
// this will be removed anyway, so no point in doing anything more than a hardcoded placeholder
fn validate_client_hello_timestamp(
client_timestamp: u64,
tolerance: Duration,
) -> Result<(), LpError> {
let now = current_timestamp()?;
let age = now.abs_diff(client_timestamp);
if age > tolerance.as_secs() {
let direction = if now >= client_timestamp {
"old"
} else {
"future"
};
return Err(LpError::kkt_psq_handshake(format!(
"ClientHello timestamp is too {direction} (age: {age}s, tolerance: {}s)",
tolerance.as_secs()
)));
}
Ok(())
pub(crate) struct PSQHandshakeStateResponder<'a, S> {
pub(super) inner_state: PSQHandshakeState<'a, S>,
pub(super) responder_data: ResponderData,
}
impl<'a, S> PSQHandshakeState<'a, S>
impl<'a, S> PSQHandshakeStateResponder<'a, S>
where
S: LpTransport + Unpin,
{
pub(crate) fn encapsulated_kem_keys(&self) -> Result<((), ()), LpError> {
todo!()
}
//
// pub(crate) fn encapsulated_kem_keys(
// &self,
// ) -> Result<(DecapsulationKey<'static>, EncapsulationKey<'static>), LpError> {
// let kem_keys = self
// .local_peer
// .kem_psq
// .as_ref()
// .ok_or(LpError::ResponderWithMissingKEMKey)?;
//
// let libcrux_private_key = libcrux_kem::PrivateKey::decode(
// libcrux_kem::Algorithm::X25519,
// kem_keys.private_key().as_bytes(),
// )
// .map_err(|e| {
// LpError::KKTError(format!(
// "Failed to convert X25519 private key to libcrux PrivateKey: {e:?}",
// ))
// })?;
// let dec_key = DecapsulationKey::X25519(libcrux_private_key);
//
// let libcrux_public_key = libcrux_kem::PublicKey::decode(
// libcrux_kem::Algorithm::X25519,
// kem_keys.public_key().as_bytes(),
// )
// .map_err(|e| {
// LpError::KKTError(format!(
// "Failed to convert X25519 public key to libcrux PublicKey: {e:?}",
// ))
// })?;
// let enc_key = EncapsulationKey::X25519(libcrux_public_key);
// Ok((dec_key, enc_key))
// }
/// Attempt to receive and validate ClientHello
pub(crate) async fn receive_client_hello(
&mut self,
) -> Result<(ClientHelloData, LpRemotePeer), LpError> {
todo!()
// let client_hello_packet = self.receive_non_error(None).await?;
// let client_hello = match client_hello_packet.message {
// LpMessage::ClientHello(client_hello) => client_hello,
// other => {
// return Err(LpError::unexpected_handshake_response(
// other.typ(),
// MessageType::ClientHello,
// ));
// }
// };
//
// validate_client_hello_timestamp(
// client_hello.extract_timestamp(),
// DEFAULT_TIMESTAMP_TOLERANCE,
// )?;
//
// // TODO: somehow check for collision
//
// // set version and remote peer information
// self.protocol_version = Some(client_hello_packet.header.protocol_version);
// let remote_peer = LpRemotePeer::new(
// client_hello.client_ed25519_public_key,
// client_hello.client_lp_public_key,
// );
//
// Ok((client_hello, remote_peer))
/// Attempt to receive a KKT request
async fn receive_kkt_request(&mut self) -> Result<KKTRequest, LpError> {
let data = self.inner_state.connection.receive_raw_packet().await?;
Ok(KKTRequest::try_from_bytes(&data)?)
}
/// Send client hello ACK
pub(crate) async fn send_client_hello_ack(&mut self, session_id: u32) -> Result<(), LpError> {
todo!()
// let protocol = self.protocol_version()?;
//
// let ack = self.next_packet(session_id, protocol, LpMessage::Ack);
// self.connection.send_packet(ack, None).await?;
// Ok(())
}
/// Attempt to process the received KKT request
fn process_kkt_request(&self, kkt_request: KKTRequest) -> Result<ProcessedKKTRequest, LpError> {
let kem_keys = &self
.inner_state
.local_peer
.kem_keypairs
.as_ref()
.ok_or(LpError::ResponderWithMissingKEMKeys)?;
/// Attempt to receive and process a KKT request
pub(crate) async fn receive_kkt_request(&mut self) -> Result<(KKTContext, (), ()), LpError> {
todo!()
let processed_req = KKTResponder::new(
&self.inner_state.local_peer.x25519,
kem_keys,
&self.responder_data.supported_hash_functions,
&self.responder_data.supported_signature_schemes,
&self.responder_data.supported_outer_protocol_versions,
)?
.process_request(kkt_request)?;
Ok(processed_req)
}
// pub(crate) async fn receive_kkt_request(
// &mut self,
// ) -> Result<(KKTContext, KKTSessionSecret, KKTSessionId), LpError> {
// let kkt_request = match self.receive_non_error(None).await?.message {
// LpMessage::KKTRequest(request) => request.0,
// other => {
// return Err(LpError::unexpected_handshake_response(
// other.typ(),
// MessageType::KKTRequest,
// ));
// }
// };
//
// let (session_secret, request_frame, remote_context) =
// decrypt_initial_kkt_frame(self.local_peer.x25519.private_key(), &kkt_request)?;
// let (context, _) = responder_ingest_message(&remote_context, None, None, &request_frame)?;
//
// Ok((context, session_secret, request_frame.session_id()))
// }
/// Attempt to send KKT response to the previously received request
pub(crate) async fn send_kkt_response(
&mut self,
session_id: u32,
// (kkt_context, session_secret, kkt_session_id): (KKTContext, KKTSessionSecret, KKTSessionId),
(kkt_context, session_secret, kkt_session_id): (KKTContext, (), ()),
// encapsulation_key: &EncapsulationKey<'_>,
encapsulation_key: &(),
) -> Result<(), LpError> {
todo!()
// let protocol = self.protocol_version()?;
//
// let response_frame = responder_process(
// &kkt_context,
// kkt_session_id,
// self.local_peer.ed25519().private_key(),
// encapsulation_key,
// )?;
// let encrypted_frame = encrypt_kkt_frame(
// &mut rng(),
// &session_secret,
// &response_frame,
// KKT_RESPONSE_AAD,
// )?;
// let lp_message = KKTResponseData::new(encrypted_frame).into();
// let lp_packet = self.next_packet(session_id, protocol, lp_message);
//
// self.connection.send_packet(lp_packet, None).await?;
// Ok(())
async fn send_kkt_response(&mut self, response: KKTResponse) -> Result<(), LpError> {
self.inner_state
.connection
.send_serialised_packet(&response.into_bytes())
.await?;
Ok(())
}
/// Attempt to receive and process a PSQ msg1 request
pub(crate) async fn receive_psq_initiator_message(
&mut self,
remote_peer: &LpRemotePeer,
// local_kem_keypair: (&DecapsulationKey<'_>, &EncapsulationKey<'_>),
local_kem_keypair: ((), ()),
salt: &[u8; 32],
session_id_bytes: &[u8; 4],
) -> Result<(OuterAeadKey, NoiseProtocol, PqSharedSecret, Vec<u8>), LpError> {
todo!()
// let psq_msg1 = match self.receive_non_error(None).await?.message {
// LpMessage::Handshake(response) => response.0,
// other => {
// return Err(LpError::unexpected_handshake_response(
// other.typ(),
// MessageType::Handshake,
// ));
// }
// };
//
// // Extract PSQ payload: [u16 psq_len][psq_payload][noise_msg]
// if psq_msg1.len() < 2 {
// return Err(LpError::kkt_psq_handshake("too short msg1 received"));
// }
// let handle_len = u16::from_le_bytes([psq_msg1[0], psq_msg1[1]]) as usize;
// if psq_msg1.len() < 2 + handle_len {
// return Err(LpError::kkt_psq_handshake("too short msg1 received"));
// }
// let psq_payload = &psq_msg1[2..2 + handle_len];
// let noise_payload = &psq_msg1[2 + handle_len..];
//
// // Decapsulate PSK from PSQ payload using X25519 as DHKEM
// let psq_responder = psq_responder_process_message(
// self.local_peer.x25519.private_key(),
// &remote_peer.x25519_public,
// local_kem_keypair,
// &remote_peer.ed25519_public,
// psq_payload,
// salt,
// session_id_bytes,
// )?;
//
// let psk = psq_responder.psk;
// let psk_handle = psq_responder.psk_handle;
//
// // TEMP \/
// let outer_aead_key = OuterAeadKey::from_psk(&psk);
// // TEMP /\
//
// let mut noise_protocol = NoiseProtocol::build_new_responder(
// self.local_peer.x25519().private_key().as_bytes(),
// remote_peer.x25519_public.as_bytes(),
// &psk,
// )?;
// noise_protocol.read_message(noise_payload)?;
//
// Ok((
// outer_aead_key,
// noise_protocol,
// PqSharedSecret::new(psq_responder.pq_shared_secret),
// psk_handle,
// ))
async fn receive_psq_initiator_message(&mut self) -> Result<Vec<u8>, LpError> {
Ok(self.inner_state.connection.receive_raw_packet().await?)
}
/// Attempt to prepare and generate a responder PSQ msg2
pub(crate) async fn send_psq_responder_message(
&mut self,
session_id: u32,
psk_handle: &[u8],
outer_aead_key: &OuterAeadKey,
noise_protocol: &mut NoiseProtocol,
) -> Result<(), LpError> {
async fn send_psq_responder_message(&mut self) -> Result<(), LpError> {
todo!()
// let protocol = self.protocol_version()?;
//
// let msg2 = noise_protocol
// .get_bytes_to_send()
// .ok_or_else(|| LpError::kkt_psq_handshake("failed to generate noise msg2"))??;
// // Embed PSK handle in message: [u16 handle_len][handle_bytes][noise_msg]
// let handle_len = psk_handle.len() as u16;
// let mut combined = Vec::with_capacity(2 + psk_handle.len() + msg2.len());
// combined.extend_from_slice(&handle_len.to_le_bytes());
// combined.extend_from_slice(psk_handle);
// combined.extend_from_slice(&msg2);
//
// let lp_message = HandshakeData::new(combined).into();
// let lp_packet = self.next_packet(session_id, protocol, lp_message);
// self.connection
// .send_packet(lp_packet, Some(outer_aead_key))
// .await?;
// Ok(())
}
/// Attempt to receive and process final PSQ msg3
pub(crate) async fn receive_final_psq_message(
&mut self,
outer_aead_key: &OuterAeadKey,
noise_protocol: &mut NoiseProtocol,
) -> Result<(), LpError> {
async fn receive_final_psq_message(&mut self) -> Result<(), LpError> {
todo!()
// let psq_msg3 = match self
// .connection
// .receive_packet(Some(outer_aead_key))
// .await?
// .message
// {
// LpMessage::Handshake(response) => response.0,
// other => {
// return Err(LpError::unexpected_handshake_response(
// other.typ(),
// MessageType::Handshake,
// ));
// }
// };
//
// noise_protocol.read_message(&psq_msg3)?;
// if !noise_protocol.is_handshake_finished() {
// return Err(LpError::kkt_psq_handshake(
// "noise handshake not finished after msg3",
// ));
// }
// Ok(())
}
/// Send final ACK to indicate finalisation of the handshake
pub(crate) async fn send_final_ack(
&mut self,
session_id: u32,
outer_aead_key: &OuterAeadKey,
) -> Result<(), LpError> {
todo!()
// let protocol = self.protocol_version()?;
//
// let ack = self.next_packet(session_id, protocol, LpMessage::Ack);
// self.connection
// .send_packet(ack, Some(outer_aead_key))
// .await?;
// Ok(())
}
async fn complete_as_responder_inner(
&mut self,
) -> Result<LpSession, IntermediateHandshakeFailure>
pub async fn complete_handshake<R>(mut self, rng: &mut R) -> Result<MinimalSession, LpError>
where
S: LpTransport + Unpin,
R: rand09::CryptoRng,
{
todo!()
// // 1. receive and validate ClientHello
// let (client_hello_data, remote_peer) =
// self.receive_client_hello()
// .await
// .map_err(|source| IntermediateHandshakeFailure {
// session_id: None,
// protocol_version: self.protocol_version,
// outer_aead_key: None,
// source,
// })?;
// debug!("received client hello");
//
// let session_id = client_hello_data.receiver_index;
// let session_id_bytes = session_id.to_le_bytes();
// let salt = client_hello_data.salt;
//
// // 2. send ack
// debug!("sending client hello ACK");
// self.send_client_hello_ack(session_id)
// .await
// .map_err(|source| IntermediateHandshakeFailure {
// session_id: Some(session_id),
// protocol_version: self.protocol_version,
// outer_aead_key: None,
// source,
// })?;
//
// // 3. receive and process KKT request
// let kkt_data =
// self.receive_kkt_request()
// .await
// .map_err(|source| IntermediateHandshakeFailure {
// session_id: Some(session_id),
// protocol_version: self.protocol_version,
// outer_aead_key: None,
// source,
// })?;
// debug!("received KKT request");
//
// todo!()
// // // TEMP: 'derive' KEM keys
// // let (dec_key, enc_key) =
// // self.encapsulated_kem_keys()
// // .map_err(|source| IntermediateHandshakeFailure {
// // session_id: Some(session_id),
// // protocol_version: self.protocol_version,
// // outer_aead_key: None,
// // source,
// // })?;
// //
// // // 4. prepare and send KKT response
// // debug!("sending KKT response");
// // self.send_kkt_response(session_id, kkt_data, &enc_key)
// // .await
// // .map_err(|source| IntermediateHandshakeFailure {
// // session_id: Some(session_id),
// // protocol_version: self.protocol_version,
// // outer_aead_key: None,
// // source,
// // })?;
// //
// // // 5. receive and process PSQ msg1
// // debug!("received PSQ msg1");
// // let (outer_aead_key, mut noise_protocol, pq_shared_secret, psk_handle) = self
// // .receive_psq_initiator_message(
// // &remote_peer,
// // (&dec_key, &enc_key),
// // &salt,
// // &session_id_bytes,
// // )
// // .await
// // .map_err(|source| IntermediateHandshakeFailure {
// // session_id: Some(session_id),
// // protocol_version: self.protocol_version,
// // outer_aead_key: None,
// // source,
// // })?;
// //
// // // 6. prepare and send PSQ msg2
// // debug!("sending PSQ msg2");
// // if let Err(source) = self
// // .send_psq_responder_message(
// // session_id,
// // &psk_handle,
// // &outer_aead_key,
// // &mut noise_protocol,
// // )
// // .await
// // {
// // return Err(IntermediateHandshakeFailure {
// // session_id: Some(session_id),
// // protocol_version: self.protocol_version,
// // outer_aead_key: Some(outer_aead_key),
// // source,
// // });
// // }
// //
// // // 7. receive and process PSQ msg3
// // debug!("received PSQ msg3");
// // if let Err(source) = self
// // .receive_final_psq_message(&outer_aead_key, &mut noise_protocol)
// // .await
// // {
// // return Err(IntermediateHandshakeFailure {
// // session_id: Some(session_id),
// // protocol_version: self.protocol_version,
// // outer_aead_key: Some(outer_aead_key),
// // source,
// // });
// // }
// //
// // // 8. [optionally] send ACK to finalise
// // debug!("sending final ACK");
// // if let Err(source) = self.send_final_ack(session_id, &outer_aead_key).await {
// // return Err(IntermediateHandshakeFailure {
// // session_id: Some(session_id),
// // protocol_version: self.protocol_version,
// // outer_aead_key: Some(outer_aead_key),
// // source,
// // });
// // }
// //
// // #[allow(clippy::expect_used)]
// // Ok(LpSession::new(
// // session_id,
// // self.protocol_version()
// // .expect("protocol version is known at this point"),
// // outer_aead_key,
// // self.local_peer.clone(),
// // remote_peer,
// // pq_shared_secret,
// // noise_protocol,
// // ))
}
// 1. receive and process KKTRequest
let kkt_request = self.receive_kkt_request().await?;
debug!("received KKT request");
pub async fn complete_as_responder(mut self) -> Result<LpSession, LpError>
where
S: LpTransport + Unpin,
{
todo!()
// match self.complete_as_responder_inner().await {
// Ok(res) => Ok(res),
// Err(err) => Err(self.try_send_error_packet(err).await),
// }
let processed_req = self.process_kkt_request(kkt_request)?;
// 2. send back the KKTResponse
debug!("sending KKT response");
self.send_kkt_response(processed_req.response).await?;
// 3. receive and process PSQ request
let raw_psq1 = self.receive_psq_initiator_message().await?;
debug!("received PSQ handshake msg");
// construct the responder and process the message
let kem = processed_req.requested_kem;
let responder_ciphersuite = build_psq_ciphersuite(&self.inner_state.local_peer, kem)?;
let version = processed_req.outer_protocol_version;
let mut psq_responder = build_psq_principal(rng, version, responder_ciphersuite)?;
psq_responder.read_message(&raw_psq1, &mut [])?;
let initiator_authenticator = psq_responder
.initiator_authenticator()
.ok_or(LpError::MissingInitiatorAuthenticator)?;
// 4. send PSQ response
let mut conn = self.inner_state.connection;
let TODO = "change buf size";
let mut buf = [0u8; 2048];
let n = psq_responder.write_message(&[], &mut buf)?;
debug!("sending PSQ handshake msg");
conn.send_serialised_packet(&buf[..n]).await?;
if !psq_responder.is_handshake_finished() {
return Err(LpError::kkt_psq_handshake(
"handshake not finished after receiving psq response",
));
}
let session = psq_responder.into_session()?;
Ok(MinimalSession {
session,
encapsulation_key: processed_req.remote_encapsulation_key,
init_authenticator: Some(initiator_authenticator),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::peer::mock_peers;
use crate::psq::initiator;
use libcrux_psq::handshake::types::Authenticator;
use libcrux_psq::session::{Session, SessionBinding};
use nym_kkt::initiator::KKTInitiator;
use nym_kkt_ciphersuite::Ciphersuite;
use nym_test_utils::helpers::{
DeterministicRng09Send, deterministic_rng_09, u64_seeded_rng_09,
};
use nym_test_utils::mocks::async_read_write::MockIOStream;
use nym_test_utils::traits::{Leak, Timeboxed};
#[tokio::test]
async fn responder_test_plain() -> anyhow::Result<()> {
let conn_init = MockIOStream::default();
let conn_resp = conn_init.try_get_remote_handle();
// SETUP START:
// leak the connections (JUST FOR THE PURPOSE OF THIS TEST!)
// so they'd get 'static lifetime
let conn_init = conn_init.leak();
let conn_resp = conn_resp.leak();
let (init, resp) = mock_peers();
let init_remote = init.as_remote();
let resp_remote = resp.as_remote();
let kem = KEM::MlKem768;
let ciphersuite = Ciphersuite::default().with_kem(kem);
let responder_data = ResponderData::default();
let handshake_resp =
PSQHandshakeState::new(conn_resp, ciphersuite, resp).as_responder(responder_data);
let mut resp_rng = DeterministicRng09Send::new(u64_seeded_rng_09(2));
let resp_fut = tokio::spawn(async move {
handshake_resp
.complete_handshake(&mut resp_rng)
.timeboxed()
.await
});
// initiator:
let mut rng = deterministic_rng_09();
let dir_hash = resp_remote.expected_kem_key_hash(init.ciphersuite)?;
// OneWay - MlKem
let (mut initiator, request) = KKTInitiator::generate_one_way_request(
&mut rng,
init.ciphersuite,
&resp_remote.x25519_public,
&dir_hash,
1,
)?;
// 1. send kkt request
conn_init
.send_serialised_packet(&request.into_bytes())
.timeboxed()
.await??;
// 2. receive KKT response
let resp = conn_init.receive_raw_packet().timeboxed().await??;
let kkt_response = KKTResponse::from_bytes(resp);
let response = initiator.process_response(kkt_response)?;
let encapsulation_key = response.encapsulation_key;
let initiator_ciphersuite =
initiator::build_psq_ciphersuite(&init, &resp_remote, &encapsulation_key)?;
let mut initiator =
initiator::build_psq_principal(rand09::rng(), 1, initiator_ciphersuite)?;
// 3. send PSQ msg1
// Send first message
let mut buf = [0u8; 2028];
let n = initiator.write_message(&[], &mut buf).unwrap();
conn_init
.send_serialised_packet(&buf[..n])
.timeboxed()
.await??;
// 4. receive PSQ msg2
let msg = conn_init.receive_raw_packet().timeboxed().await??;
initiator.read_message(&msg, &mut []).unwrap();
assert!(initiator.is_handshake_finished());
let session_resp = resp_fut.await???;
let init_auth = session_resp.init_authenticator.unwrap();
let i_transport = initiator.into_session().unwrap();
let r_transport = session_resp.session;
// test serialization, deserialization
let mut msg_channel = vec![0u8; 2048];
let mut payload_buf_responder = vec![0u8; 4096];
let mut payload_buf_initiator = vec![0u8; 4096];
let mut session_storage = vec![0u8; 4096];
i_transport
.serialize(
&mut session_storage,
SessionBinding {
initiator_authenticator: &Authenticator::Dh(init.x25519().pk),
responder_ecdh_pk: &resp_remote.x25519_public,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
let mut i_transport = Session::deserialize(
&session_storage,
SessionBinding {
initiator_authenticator: &Authenticator::Dh(init.x25519().pk),
responder_ecdh_pk: &resp_remote.x25519_public,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
r_transport
.serialize(
&mut session_storage,
SessionBinding {
initiator_authenticator: &init_auth,
responder_ecdh_pk: &resp_remote.x25519_public,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
let mut r_transport = Session::deserialize(
&session_storage,
SessionBinding {
initiator_authenticator: &init_auth,
responder_ecdh_pk: &resp_remote.x25519_public,
responder_pq_pk: Some(encapsulation_key.as_pq_encapsulation_key()),
},
)
.unwrap();
let mut channel_i = i_transport.transport_channel().unwrap();
let mut channel_r = r_transport.transport_channel().unwrap();
assert_eq!(channel_i.identifier(), channel_r.identifier());
let app_data_i = b"Derived session hey".as_slice();
let app_data_r = b"Derived session ho".as_slice();
let len_i = channel_i
.write_message(app_data_i, &mut msg_channel)
.unwrap();
let (len_r_deserialized, len_r_payload) = channel_r
.read_message(&msg_channel, &mut payload_buf_responder)
.unwrap();
// We read the same amount of data.
assert_eq!(len_r_deserialized, len_i);
assert_eq!(len_r_payload, app_data_i.len());
assert_eq!(&payload_buf_responder[0..len_r_payload], app_data_i);
let len_r = channel_r
.write_message(app_data_r, &mut msg_channel)
.unwrap();
let (len_i_deserialized, len_i_payload) = channel_i
.read_message(&msg_channel, &mut payload_buf_initiator)
.unwrap();
assert_eq!(len_r, len_i_deserialized);
assert_eq!(app_data_r.len(), len_i_payload);
assert_eq!(&payload_buf_initiator[0..len_i_payload], app_data_r);
Ok(())
}
}
+7 -5
View File
@@ -158,7 +158,7 @@ impl LpSession {
}
/// Helper function to create `PSQHandshakeState` for the handshake initiator
pub fn complete_as_initiator<S>(
pub fn psq_handshake_initiator<S>(
connection: &'_ mut S,
ciphersuite: Ciphersuite,
local_peer: LpLocalPeer,
@@ -168,9 +168,10 @@ impl LpSession {
where
S: LpTransport + Unpin,
{
PSQHandshakeState::new(connection, ciphersuite, local_peer)
.with_protocol_version(remote_protocol_version)
.with_remote_peer(remote_peer)
todo!()
// PSQHandshakeState::new(connection, ciphersuite, local_peer)
// .with_protocol_version(remote_protocol_version)
// .with_remote_peer(remote_peer)
}
/// Helper function to create `PSQHandshakeState` for the handshake responder
@@ -182,7 +183,8 @@ impl LpSession {
where
S: LpTransport + Unpin,
{
PSQHandshakeState::new(connection, ciphersuite, local_peer)
todo!()
// PSQHandshakeState::new(connection, ciphersuite, local_peer)
}
pub fn id(&self) -> u32 {
@@ -414,7 +414,7 @@ where
// TODO:
let ciphersuite = LpSession::default_ciphersuite();
let session = LpSession::complete_as_initiator(
let session = LpSession::psq_handshake_initiator(
connection,
ciphersuite,
local_peer,
@@ -201,7 +201,7 @@ impl NestedLpSession {
let protocol_version = self.gateway_supported_lp_protocol_version;
let ciphersuite = LpSession::default_ciphersuite();
let session = LpSession::complete_as_initiator(
let session = LpSession::psq_handshake_initiator(
&mut nested_connection,
ciphersuite,
local_peer,