pub(super) mod proto {
#![allow(unreachable_pub)]
include!("../generated/mod.rs");
pub use self::payload::proto::{NoiseExtensions, NoiseHandshakePayload};
}
use std::{collections::HashSet, io, mem};
use asynchronous_codec::Framed;
use futures::prelude::*;
use libp2p_identity as identity;
use multihash::Multihash;
use quick_protobuf::MessageWrite;
use super::framed::Codec;
use crate::{
io::Output,
protocol::{KeypairIdentity, PublicKey, STATIC_KEY_DOMAIN},
Error,
};
pub(crate) struct State<T> {
io: Framed<T, Codec<snow::HandshakeState>>,
identity: KeypairIdentity,
dh_remote_pubkey_sig: Option<Vec<u8>>,
id_remote_pubkey: Option<identity::PublicKey>,
responder_webtransport_certhashes: Option<HashSet<Multihash<64>>>,
remote_extensions: Option<Extensions>,
}
struct Extensions {
webtransport_certhashes: HashSet<Multihash<64>>,
}
impl<T> State<T>
where
T: AsyncRead + AsyncWrite,
{
pub(crate) fn new(
io: T,
session: snow::HandshakeState,
identity: KeypairIdentity,
expected_remote_key: Option<identity::PublicKey>,
responder_webtransport_certhashes: Option<HashSet<Multihash<64>>>,
) -> Self {
Self {
identity,
io: Framed::new(io, Codec::new(session)),
dh_remote_pubkey_sig: None,
id_remote_pubkey: expected_remote_key,
responder_webtransport_certhashes,
remote_extensions: None,
}
}
}
impl<T> State<T>
where
T: AsyncRead + AsyncWrite,
{
pub(crate) fn finish(self) -> Result<(identity::PublicKey, Output<T>), Error> {
let is_initiator = self.io.codec().is_initiator();
let (pubkey, framed) = map_into_transport(self.io)?;
let id_pk = self
.id_remote_pubkey
.ok_or_else(|| Error::AuthenticationFailed)?;
let is_valid_signature = self.dh_remote_pubkey_sig.as_ref().is_some_and(|s| {
id_pk.verify(&[STATIC_KEY_DOMAIN.as_bytes(), pubkey.as_ref()].concat(), s)
});
if !is_valid_signature {
return Err(Error::BadSignature);
}
if is_initiator {
if let Some(expected_certhashes) = self.responder_webtransport_certhashes {
let ext = self.remote_extensions.ok_or_else(|| {
Error::UnknownWebTransportCerthashes(
expected_certhashes.to_owned(),
HashSet::new(),
)
})?;
let received_certhashes = ext.webtransport_certhashes;
if !expected_certhashes.is_subset(&received_certhashes) {
return Err(Error::UnknownWebTransportCerthashes(
expected_certhashes,
received_certhashes,
));
}
}
}
Ok((id_pk, Output::new(framed)))
}
}
fn map_into_transport<T>(
framed: Framed<T, Codec<snow::HandshakeState>>,
) -> Result<(PublicKey, Framed<T, Codec<snow::TransportState>>), Error>
where
T: AsyncRead + AsyncWrite,
{
let mut parts = framed.into_parts().map_codec(Some);
let (pubkey, codec) = mem::take(&mut parts.codec)
.expect("We just set it to `Some`")
.into_transport()?;
let parts = parts.map_codec(|_| codec);
let framed = Framed::from_parts(parts);
Ok((pubkey, framed))
}
impl From<proto::NoiseExtensions> for Extensions {
fn from(value: proto::NoiseExtensions) -> Self {
Extensions {
webtransport_certhashes: value
.webtransport_certhashes
.into_iter()
.filter_map(|bytes| Multihash::read(&bytes[..]).ok())
.collect(),
}
}
}
async fn recv<T>(state: &mut State<T>) -> Result<proto::NoiseHandshakePayload, Error>
where
T: AsyncRead + Unpin,
{
match state.io.next().await {
None => Err(io::Error::new(io::ErrorKind::UnexpectedEof, "eof").into()),
Some(Err(e)) => Err(e.into()),
Some(Ok(p)) => Ok(p),
}
}
pub(crate) async fn recv_empty<T>(state: &mut State<T>) -> Result<(), Error>
where
T: AsyncRead + Unpin,
{
let payload = recv(state).await?;
if payload.get_size() != 0 {
return Err(io::Error::new(io::ErrorKind::InvalidData, "Expected empty payload.").into());
}
Ok(())
}
pub(crate) async fn send_empty<T>(state: &mut State<T>) -> Result<(), Error>
where
T: AsyncWrite + Unpin,
{
state
.io
.send(&proto::NoiseHandshakePayload::default())
.await?;
Ok(())
}
pub(crate) async fn recv_identity<T>(state: &mut State<T>) -> Result<(), Error>
where
T: AsyncRead + Unpin,
{
let pb = recv(state).await?;
state.id_remote_pubkey = Some(identity::PublicKey::try_decode_protobuf(&pb.identity_key)?);
if !pb.identity_sig.is_empty() {
state.dh_remote_pubkey_sig = Some(pb.identity_sig);
}
if let Some(extensions) = pb.extensions {
state.remote_extensions = Some(extensions.into());
}
Ok(())
}
pub(crate) async fn send_identity<T>(state: &mut State<T>) -> Result<(), Error>
where
T: AsyncRead + AsyncWrite + Unpin,
{
let mut pb = proto::NoiseHandshakePayload {
identity_key: state.identity.public.encode_protobuf(),
..Default::default()
};
pb.identity_sig.clone_from(&state.identity.signature);
if state.io.codec().is_responder() {
if let Some(ref certhashes) = state.responder_webtransport_certhashes {
let ext = pb
.extensions
.get_or_insert_with(proto::NoiseExtensions::default);
ext.webtransport_certhashes = certhashes.iter().map(|hash| hash.to_bytes()).collect();
}
}
state.io.send(&pb).await?;
Ok(())
}