[WIP] Security work While on holiday

This commit is contained in:
Alex 2026-08-12 22:45:28 +02:00
commit 7f0231e3f1
Signed by: alex
SSH key fingerprint: SHA256:D1+Ub8o0v4K5y1JNivW8IxEOelqLSvPmUzBbDIoZkRQ
109 changed files with 19694 additions and 5210 deletions

View file

@ -7,11 +7,15 @@
use crate::config::HostConfig;
use crate::error::AcceptError;
use mtp_codec::{
CommunicationType, CommunicationValue, DataType, DataValue, Version,
CommunicationType, CommunicationValue, DataType, DataValue, TypeMap, Version,
registry::{Registry, VersionedCodec},
};
use mtp_common::{CommunicationError, RejectionReason};
#[cfg(feature = "crypto")]
use std::collections::HashSet;
use std::sync::Arc;
#[cfg(feature = "crypto")]
use std::sync::Mutex;
/// Trait for sending handshake messages during the opening exchange.
///
@ -24,6 +28,9 @@ pub trait HandshakeSender: Send + Sync {
fn finish_stream(
&self,
) -> impl std::future::Future<Output = Result<(), CommunicationError>> + Send;
fn set_type_map(&self, _type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async {}
}
fn close(&self);
}
@ -34,6 +41,37 @@ pub trait HandshakeReceiver: Send + Sync {
fn receive(
&self,
) -> impl std::future::Future<Output = Result<CommunicationValue, CommunicationError>> + Send;
/// Bind subsequently decoded frames to the negotiated type map.
///
/// The opening frame must be decoded with the transport's bootstrap map so
/// that it can reveal the version. Once negotiation succeeds, all later
/// frames—including the remainder of the authentication exchange—must use
/// the negotiated map rather than whichever map happens to be latest at
/// compile time.
fn set_type_map(&self, _type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async {}
}
}
/// A non-zero guest ID reserved for the lifetime of a connected session.
///
/// The lease is moved into the resulting `MTPConnection`, so dropping that
/// connection releases the ID for a later guest session.
#[cfg(feature = "crypto")]
#[derive(Debug)]
pub struct GuestIdLease {
active_ids: Arc<Mutex<HashSet<u64>>>,
id: u64,
}
#[cfg(feature = "crypto")]
impl Drop for GuestIdLease {
fn drop(&mut self) {
if let Ok(mut active_ids) = self.active_ids.lock() {
active_ids.remove(&self.id);
}
}
}
/// The result of a successful handshake, containing everything needed to
@ -49,6 +87,8 @@ pub struct HandshakeResult {
pub client_id: u64,
#[cfg(feature = "crypto")]
pub client_public_key: Option<mtp_crypto::PublicKeyBundle>,
#[cfg(feature = "crypto")]
pub guest_id_lease: Option<GuestIdLease>,
}
/// Transport-independent handshake state machine.
@ -90,13 +130,56 @@ impl HandshakeEngine {
) -> Result<HandshakeResult, AcceptError> {
#[cfg(feature = "crypto")]
{
let timeout = self.config.auth_timeout;
tokio::time::timeout(timeout, self.accept_inner(sender, receiver))
.await
.unwrap_or(Err(AcceptError::AuthenticationTimedOut))
self.accept_until(
sender,
receiver,
tokio::time::Instant::now() + self.config.auth_timeout,
)
.await
}
#[cfg(not(feature = "crypto"))]
self.accept_inner(sender, receiver).await
{
let result = self.accept_inner(sender, receiver).await;
if result.is_err() {
sender.close();
}
result
}
}
/// Run the crypto handshake until an absolute deadline.
///
/// WebTransport authentication may wait for a shared semaphore before it
/// reaches this engine. Passing the deadline through keeps that queueing
/// time from silently granting the handshake another full timeout.
#[cfg(feature = "crypto")]
pub async fn accept_until<S: HandshakeSender, R: HandshakeReceiver>(
&self,
sender: &S,
receiver: &R,
deadline: tokio::time::Instant,
) -> Result<HandshakeResult, AcceptError> {
match tokio::time::timeout_at(deadline, self.accept_inner(sender, receiver)).await {
Ok(result) => {
if result.is_err() {
sender.close();
}
result
}
Err(_) => {
let error = AcceptError::AuthenticationTimedOut;
send_rejection_generic(
sender,
RejectionReason::AuthenticationFailed {
detail: error.to_string(),
},
None,
)
.await;
sender.close();
Err(error)
}
}
}
async fn accept_inner<S: HandshakeSender, R: HandshakeReceiver>(
@ -104,16 +187,17 @@ impl HandshakeEngine {
sender: &S,
receiver: &R,
) -> Result<HandshakeResult, AcceptError> {
let first_msg = receiver.receive().await.map_err(AcceptError::Receive)?;
let mut first_msg = receiver.receive().await.map_err(AcceptError::Receive)?;
let version_str = match first_msg.get_data(DataType::Version) {
DataValue::Str(s) => s.clone(),
Some(DataValue::Str(s)) => s.clone(),
_ => {
send_rejection_generic(
sender,
RejectionReason::AuthenticationFailed {
detail: "opening message omitted a valid protocol version".into(),
},
None,
)
.await;
sender.close();
@ -128,6 +212,7 @@ impl HandshakeEngine {
RejectionReason::AuthenticationFailed {
detail: "opening message omitted a valid protocol version".into(),
},
None,
)
.await;
sender.close();
@ -150,6 +235,7 @@ impl HandshakeEngine {
.map(|v| v.to_string())
.collect(),
},
None,
)
.await;
sender.close();
@ -160,8 +246,12 @@ impl HandshakeEngine {
let codec = VersionedCodec::for_version(self.registry.clone(), negotiated.clone())
.ok_or_else(|| AcceptError::UnsupportedVersion(negotiated.clone()))?;
sender.set_type_map(codec.type_map()).await;
receiver.set_type_map(codec.type_map()).await;
first_msg.set_type_map(codec.type_map());
let description = match first_msg.get_data(DataType::Description) {
DataValue::Str(s) => Some(s.clone()),
Some(DataValue::Str(s)) => Some(s.clone()),
_ => None,
};
@ -209,9 +299,28 @@ impl HandshakeEngine {
#[cfg(not(feature = "crypto"))]
{
let _ = sender;
let authentication_requested = Some(first_msg.get_type())
== CommunicationType::Register.try_to_id(codec.type_map())
|| first_msg.get_data(DataType::PublicKeys).is_some();
if authentication_requested {
let error = AcceptError::AuthenticationFailed(
"authentication is unavailable on this non-crypto host".into(),
);
send_rejection_generic(
sender,
RejectionReason::AuthenticationFailed {
detail: error.to_string(),
},
Some(codec.type_map()),
)
.await;
sender.close();
return Err(error);
}
let _ = receiver;
let _ = first_msg;
send_accepted_generic(sender, &negotiated, codec.type_map(), Some(0))
.await
.map_err(AcceptError::Send)?;
Ok(HandshakeResult {
negotiated_version: negotiated,
codec,
@ -229,15 +338,21 @@ impl HandshakeEngine {
codec: VersionedCodec,
description: Option<String>,
) -> Result<HandshakeResult, AcceptError> {
let tm = mtp_codec::TypeMap::latest();
let tm = codec.type_map();
// Reject Register frames on unauthenticated hosts
if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(&tm) {
// Reject explicit authentication attempts on unauthenticated hosts.
// Authenticated clients include PublicKeys in Identification as an
// intent marker; this avoids acknowledging the opening as a guest
// connection and leaving the client waiting for a Challenge.
if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(&tm)
|| first_msg.get_data(DataType::PublicKeys).is_some()
{
send_rejection_generic(
sender,
RejectionReason::AuthenticationFailed {
detail: "authentication not allowed on this host".into(),
},
Some(tm),
)
.await;
sender.close();
@ -246,8 +361,15 @@ impl HandshakeEngine {
));
}
let guest_id = self.assign_guest_id().await?;
send_accepted_generic(sender, &negotiated, Some(guest_id))
let guest_id_lease = match self.assign_guest_id().await {
Ok(lease) => lease,
Err(error) => {
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
};
let guest_id = guest_id_lease.id;
send_accepted_generic(sender, &negotiated, tm, Some(guest_id))
.await
.map_err(AcceptError::Send)?;
@ -258,6 +380,7 @@ impl HandshakeEngine {
auth_state: crate::error::AuthState::Unauthenticated,
client_id: guest_id,
client_public_key: None,
guest_id_lease: Some(guest_id_lease),
})
}
@ -274,11 +397,17 @@ impl HandshakeEngine {
version_str: &str,
client_version: Version,
) -> Result<HandshakeResult, AcceptError> {
let tm = mtp_codec::TypeMap::latest();
let tm = codec.type_map();
// Register frames always go through full authentication
if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(&tm) {
let bundle = extract_register_bundle(&first_msg)?;
let bundle = match extract_register_bundle(&first_msg) {
Ok(bundle) => bundle,
Err(error) => {
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
};
let pk_bytes = bundle.as_bytes();
return self
.complete_auth_handshake(
@ -298,7 +427,7 @@ impl HandshakeEngine {
// Identification: try lookup, fall back to guest
if Some(first_msg.get_type()) == CommunicationType::Identification.try_to_id(&tm) {
let cid = match first_msg.get_data(DataType::Id) {
DataValue::UnsignedNumber(n) => *n as u64,
Some(DataValue::UnsignedNumber(n)) => u64::try_from(*n).unwrap_or(0),
_ => 0,
};
@ -321,9 +450,26 @@ impl HandshakeEngine {
.await;
}
// Unknown or zero ID: an Identification carrying PublicKeys is an
// explicit authentication attempt, not a guest connection.
if first_msg.get_data(DataType::PublicKeys).is_some() {
let error = AcceptError::AuthenticationFailed(
"unknown authenticated client identity".into(),
);
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
// Unknown or zero ID: fall back to guest
let guest_id = self.assign_guest_id().await?;
send_accepted_generic(sender, &negotiated, Some(guest_id))
let guest_id_lease = match self.assign_guest_id().await {
Ok(lease) => lease,
Err(error) => {
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
};
let guest_id = guest_id_lease.id;
send_accepted_generic(sender, &negotiated, tm, Some(guest_id))
.await
.map_err(AcceptError::Send)?;
return Ok(HandshakeResult {
@ -333,13 +479,13 @@ impl HandshakeEngine {
auth_state: crate::error::AuthState::Unauthenticated,
client_id: guest_id,
client_public_key: None,
guest_id_lease: Some(guest_id_lease),
});
}
sender.close();
Err(AcceptError::AuthenticationFailed(
"unexpected message type".into(),
))
let error = AcceptError::AuthenticationFailed("unexpected message type".into());
reject_error_generic(sender, &error, tm).await;
Err(error)
}
#[cfg(feature = "crypto")]
@ -355,30 +501,39 @@ impl HandshakeEngine {
version_str: &str,
client_version: Version,
) -> Result<HandshakeResult, AcceptError> {
let tm = mtp_codec::TypeMap::latest();
let tm = codec.type_map();
let (flow, response_type) = if Some(first_msg.get_type())
== CommunicationType::Identification.try_to_id(&tm)
{
let cid = match first_msg.get_data(DataType::Id) {
DataValue::UnsignedNumber(n) => *n as u64,
Some(DataValue::UnsignedNumber(n)) => match u64::try_from(*n) {
Ok(id) => id,
Err(_) => {
let error =
AcceptError::AuthenticationFailed("client id is out of range".into());
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
},
_ => {
sender.close();
return Err(AcceptError::AuthenticationFailed(
"missing client id".into(),
));
let error = AcceptError::AuthenticationFailed("missing client id".into());
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
};
let bundle = match (self.config.get_existing_client)(cid, description.clone()).await {
Some(b) => b,
None => {
let rejection =
CommunicationValue::new(CommunicationType::IdentificationResponse)
.add_typed_default(DataType::Connected, DataValue::BoolFalse)
.add_typed_default(
DataType::ErrorMessage,
DataValue::Str("unknown client id".into()),
);
let rejection = CommunicationValue::new_with_type_map(
CommunicationType::IdentificationResponse,
tm,
)
.add_typed_default(DataType::Connected, DataValue::BoolFalse)
.add_typed_default(
DataType::ErrorMessage,
DataValue::Str("unknown client id".into()),
);
let _ = sender.send(&rejection).await;
sender.close();
return Err(AcceptError::AuthenticationFailed(
@ -391,17 +546,23 @@ impl HandshakeEngine {
CommunicationType::IdentificationResponse,
)
} else if Some(first_msg.get_type()) == CommunicationType::Register.try_to_id(&tm) {
let bundle = extract_register_bundle(&first_msg)?;
let bundle = match extract_register_bundle(&first_msg) {
Ok(bundle) => bundle,
Err(error) => {
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
};
let pk_bytes = bundle.as_bytes();
(
Flow::Register { bundle, pk_bytes },
CommunicationType::RegisterResponse,
)
} else {
sender.close();
return Err(AcceptError::AuthenticationFailed(
"unexpected authentication message".into(),
));
let error =
AcceptError::AuthenticationFailed("unexpected authentication message".into());
reject_error_generic(sender, &error, tm).await;
return Err(error);
};
self.complete_auth_handshake(
@ -434,7 +595,7 @@ impl HandshakeEngine {
) -> Result<HandshakeResult, AcceptError> {
use mtp_crypto::{Ed25519Signer, MlDsaSigner, SignatureScheme, auth, verify_ed25519};
let tm = mtp_codec::TypeMap::latest();
let tm = codec.type_map();
// PQ preflight: host requiring PQ must have a PQ key
let pq_enabled = !self
@ -457,6 +618,7 @@ impl HandshakeEngine {
RejectionReason::AuthenticationFailed {
detail: "host requires PQ authentication but has no PQ signing key".into(),
},
Some(tm),
)
.await;
sender.close();
@ -468,11 +630,17 @@ impl HandshakeEngine {
// Initialize host signers
let host_pq_signer = if pq_enabled {
Some(Arc::new(
MlDsaSigner::new(
match MlDsaSigner::new(
&self.config.host_keyring.sig_pq_secret_key,
&self.config.host_keyring.sig_pq_public_key,
)
.map_err(|e| AcceptError::AuthenticationFailed(e.to_string()))?,
) {
Ok(signer) => signer,
Err(error) => {
let error = AcceptError::AuthenticationFailed(error.to_string());
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
},
))
} else {
None
@ -505,14 +673,21 @@ impl HandshakeEngine {
let server_challenge: u128 = rand::random();
let (chal_sig, chal_pq_sig) =
host_sign(auth::challenge_payload(challenge_id, server_challenge)).await?;
match host_sign(auth::challenge_payload(challenge_id, server_challenge)).await {
Ok(signatures) => signatures,
Err(error) => {
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
};
let mut challenge_msg = CommunicationValue::new(CommunicationType::Challenge)
.add_typed_default(
DataType::ServerNonce,
DataValue::UnsignedNumber(server_challenge),
)
.add_typed_default(DataType::Signature, DataValue::Bytes(chal_sig));
let mut challenge_msg =
CommunicationValue::new_with_type_map(CommunicationType::Challenge, tm)
.add_typed_default(
DataType::ServerNonce,
DataValue::UnsignedNumber(server_challenge),
)
.add_typed_default(DataType::Signature, DataValue::Bytes(chal_sig));
challenge_msg = challenge_msg.add_typed_default(
DataType::RequirePq,
if self.config.require_pq {
@ -536,31 +711,28 @@ impl HandshakeEngine {
AcceptError::Receive(e)
})?;
if Some(proof.get_type()) != CommunicationType::ChallengeResponse.try_to_id(&tm) {
sender.close();
return Err(AcceptError::AuthenticationFailed(
"missing challenge response".into(),
));
let error = AcceptError::AuthenticationFailed("missing challenge response".into());
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
let client_nonce = match proof.get_data(DataType::ClientNonce) {
DataValue::UnsignedNumber(n) => *n,
Some(DataValue::UnsignedNumber(n)) => *n,
_ => {
sender.close();
return Err(AcceptError::AuthenticationFailed(
"missing client nonce".into(),
));
let error = AcceptError::AuthenticationFailed("missing client nonce".into());
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
};
let sig_bytes = match proof.get_data(DataType::Signature) {
DataValue::Bytes(b) => b.clone(),
Some(DataValue::Bytes(b)) => b.clone(),
_ => {
sender.close();
return Err(AcceptError::AuthenticationFailed(
"missing challenge signature".into(),
));
let error = AcceptError::AuthenticationFailed("missing challenge signature".into());
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
};
let pq_sig_bytes: Vec<u8> = match proof.get_data(DataType::PqSignature) {
DataValue::Bytes(b) => b.clone(),
Some(DataValue::Bytes(b)) => b.clone(),
_ => vec![],
};
@ -601,6 +773,7 @@ impl HandshakeEngine {
RejectionReason::AuthenticationFailed {
detail: "client proof signature invalid".into(),
},
Some(tm),
)
.await;
sender.close();
@ -613,21 +786,66 @@ impl HandshakeEngine {
let (assigned_id, client_bundle) = match flow {
Flow::Login { id, bundle } => (id, bundle),
Flow::Register { bundle, .. } => {
let new_id =
(self.config.complete_register)(bundle.clone(), description.clone()).await;
let _registration_guard = self.config.registration_lock.lock().await;
let identity = bundle.as_bytes();
let cached_id = self
.config
.registration_ids
.lock()
.ok()
.and_then(|registrations| registrations.get(&identity).copied());
let new_id = if let Some(id) = cached_id {
id
} else if let Some(lookup) = &self.config.find_registered_client {
match lookup(bundle.clone(), description.clone()).await {
Some(id) => id,
None => {
(self.config.complete_register)(bundle.clone(), description.clone())
.await
}
}
} else {
(self.config.complete_register)(bundle.clone(), description.clone()).await
};
if new_id != 0
&& let Ok(mut registrations) = self.config.registration_ids.lock()
{
registrations.insert(identity, new_id);
}
if new_id == 0 {
let error = AcceptError::AuthenticationFailed(
"registration callback returned reserved client id 0".into(),
);
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
(new_id, bundle)
}
};
if assigned_id == 0 {
let error = AcceptError::AuthenticationFailed(
"client id 0 is reserved for no authenticated identity".into(),
);
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
// Sign and send final response
let (host_sig, host_pq_sig) = host_sign(auth::host_final_payload(
let (host_sig, host_pq_sig) = match host_sign(auth::host_final_payload(
assigned_id,
client_nonce,
server_challenge,
))
.await?;
.await
{
Ok(signatures) => signatures,
Err(error) => {
reject_error_generic(sender, &error, tm).await;
return Err(error);
}
};
let mut response = CommunicationValue::new(response_type)
let mut response = CommunicationValue::new_with_type_map(response_type, tm)
.add_typed_default(DataType::Connected, DataValue::BoolTrue)
.add_typed_default(DataType::Id, DataValue::UnsignedNumber(assigned_id as u128))
.add_typed_default(
@ -658,6 +876,7 @@ impl HandshakeEngine {
auth_state: crate::error::AuthState::Authenticated,
client_id: assigned_id,
client_public_key: Some(client_bundle),
guest_id_lease: None,
})
}
}
@ -682,36 +901,52 @@ enum Flow {
impl HandshakeEngine {
const GUEST_ID_MAX_RETRIES: u32 = 100;
async fn assign_guest_id(&self) -> Result<u64, AcceptError> {
async fn assign_guest_id(&self) -> Result<GuestIdLease, AcceptError> {
if let Some(ref generator) = self.config.guest_id_generator {
let id = generator().await.ok_or_else(|| {
AcceptError::AuthenticationFailed(
"guest id generator rejected the connection".into(),
)
})?;
if id > mtp_codec::MAX_WIRE_ID {
if id == 0 {
return Err(AcceptError::AuthenticationFailed(
"guest id exceeds wire limit".into(),
"guest id 0 is reserved for no authenticated identity".into(),
));
}
if (self.config.get_existing_client)(id, None).await.is_none() {
return Ok(id);
if let Some(lease) = self.try_reserve_guest_id(id).await? {
return Ok(lease);
}
}
self.random_guest_id().await
}
async fn random_guest_id(&self) -> Result<u64, AcceptError> {
async fn random_guest_id(&self) -> Result<GuestIdLease, AcceptError> {
for _ in 0..Self::GUEST_ID_MAX_RETRIES {
let id = rand::random::<u64>() & mtp_codec::MAX_WIRE_ID;
if (self.config.get_existing_client)(id, None).await.is_none() {
return Ok(id);
let id = rand::random::<u64>();
if let Some(lease) = self.try_reserve_guest_id(id).await? {
return Ok(lease);
}
}
Err(AcceptError::AuthenticationFailed(
"failed to allocate a unique guest id after retries".into(),
))
}
async fn try_reserve_guest_id(&self, id: u64) -> Result<Option<GuestIdLease>, AcceptError> {
if id == 0 || (self.config.get_existing_client)(id, None).await.is_some() {
return Ok(None);
}
let mut active_ids = self.config.active_guest_ids.lock().map_err(|_| {
AcceptError::AuthenticationFailed("guest ID registry is poisoned".into())
})?;
if !active_ids.insert(id) {
return Ok(None);
}
Ok(Some(GuestIdLease {
active_ids: Arc::clone(&self.config.active_guest_ids),
id,
}))
}
}
// ---------------------------------------------------------------------------
@ -723,7 +958,7 @@ fn extract_register_bundle(
msg: &CommunicationValue,
) -> Result<mtp_crypto::PublicKeyBundle, AcceptError> {
match msg.get_data(DataType::PublicKeys) {
DataValue::Bytes(b) => mtp_crypto::PublicKeyBundle::from_bytes(b)
Some(DataValue::Bytes(b)) => mtp_crypto::PublicKeyBundle::from_bytes(b)
.map_err(|_| AcceptError::AuthenticationFailed("invalid public key bundle".into())),
_ => Err(AcceptError::AuthenticationFailed(
"missing public keys".into(),
@ -731,32 +966,58 @@ fn extract_register_bundle(
}
}
async fn send_rejection_generic<S: HandshakeSender>(sender: &S, reason: RejectionReason) {
async fn send_rejection_generic<S: HandshakeSender>(
sender: &S,
reason: RejectionReason,
type_map: Option<&TypeMap>,
) {
let type_map = type_map.cloned().unwrap_or_else(TypeMap::latest);
let response = match &reason {
RejectionReason::BadVersion { supported_versions } => {
CommunicationValue::new(CommunicationType::ErrorBadVersion)
CommunicationValue::new_with_type_map(CommunicationType::ErrorBadVersion, &type_map)
.add_typed_default(
DataType::Version,
DataValue::Str(supported_versions.join(",")),
)
.add_typed_default(DataType::ErrorMessage, DataValue::Str(reason.to_string()))
}
_ => CommunicationValue::new(CommunicationType::IdentificationResponse)
.add_typed_default(DataType::Connected, DataValue::BoolFalse)
.add_typed_default(DataType::ErrorMessage, DataValue::Str(reason.to_string())),
_ => CommunicationValue::new_with_type_map(
CommunicationType::IdentificationResponse,
&type_map,
)
.add_typed_default(DataType::Connected, DataValue::BoolFalse)
.add_typed_default(DataType::ErrorMessage, DataValue::Str(reason.to_string())),
};
let _ = sender.send(&response).await;
}
#[cfg(feature = "crypto")]
async fn reject_error_generic<S: HandshakeSender>(
sender: &S,
error: &AcceptError,
type_map: &TypeMap,
) {
send_rejection_generic(
sender,
RejectionReason::AuthenticationFailed {
detail: error.to_string(),
},
Some(type_map),
)
.await;
sender.close();
}
async fn send_accepted_generic<S: HandshakeSender>(
sender: &S,
version: &Version,
type_map: &TypeMap,
assigned_id: Option<u64>,
) -> Result<(), CommunicationError> {
let mut response = CommunicationValue::new(CommunicationType::IdentificationResponse)
.add_typed_default(DataType::Connected, DataValue::BoolTrue)
.add_typed_default(DataType::Version, DataValue::Str(version.to_string()));
let mut response =
CommunicationValue::new_with_type_map(CommunicationType::IdentificationResponse, type_map)
.add_typed_default(DataType::Connected, DataValue::BoolTrue)
.add_typed_default(DataType::Version, DataValue::Str(version.to_string()));
if let Some(id) = assigned_id {
response = response.add_typed_default(DataType::Id, DataValue::UnsignedNumber(id as u128));
}
@ -780,6 +1041,9 @@ impl HandshakeSender for mtp_transport::Sender {
) -> impl std::future::Future<Output = Result<(), CommunicationError>> + Send {
mtp_transport::Sender::finish_stream(self)
}
fn set_type_map(&self, type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async move { self.set_type_map(type_map).await }
}
fn close(&self) {
let sender = self.clone();
tokio::spawn(async move { sender.close().await });
@ -793,6 +1057,10 @@ impl HandshakeReceiver for mtp_transport::Receiver {
{
mtp_transport::Receiver::receive(self)
}
fn set_type_map(&self, type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async move { self.set_type_map(type_map).await }
}
}
impl<C: mtp_transport::TransportConnection> HandshakeSender for mtp_transport::GenericSender<C> {
@ -807,6 +1075,9 @@ impl<C: mtp_transport::TransportConnection> HandshakeSender for mtp_transport::G
) -> impl std::future::Future<Output = Result<(), CommunicationError>> + Send {
mtp_transport::GenericSender::finish_stream(self)
}
fn set_type_map(&self, type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async move { self.set_type_map(type_map).await }
}
fn close(&self) {
mtp_transport::GenericSender::close(self);
}
@ -821,4 +1092,68 @@ impl<C: mtp_transport::TransportConnection> HandshakeReceiver
{
mtp_transport::GenericReceiver::receive(self)
}
fn set_type_map(&self, type_map: &TypeMap) -> impl std::future::Future<Output = ()> + Send {
async move { self.set_type_map(type_map).await }
}
}
#[cfg(all(test, feature = "crypto"))]
mod tests {
use super::*;
use std::sync::Mutex;
#[tokio::test]
async fn guest_generator_accepts_full_width_id_after_collision_check()
-> Result<(), Box<dyn std::error::Error>> {
let lookups = Arc::new(Mutex::new(Vec::new()));
let recorded_lookups = Arc::clone(&lookups);
let mut config = HostConfig::new("127.0.0.1".parse()?, 4433, Vec::new(), Vec::new())
.with_guest_id_generator(Box::new(|| Box::pin(async { Some(u64::MAX) })));
config.get_existing_client = Box::new(move |id, description| {
let recorded_lookups = Arc::clone(&recorded_lookups);
Box::pin(async move {
recorded_lookups
.lock()
.expect("guest ID lookup mutex should not be poisoned")
.push((id, description));
None
})
});
let engine = HandshakeEngine::new(Registry::builtin(), Arc::new(config));
assert_eq!(engine.assign_guest_id().await?.id, u64::MAX);
assert_eq!(
*lookups
.lock()
.expect("guest ID lookup mutex should not be poisoned"),
vec![(u64::MAX, None)]
);
Ok(())
}
#[tokio::test]
async fn active_guest_ids_are_unique_and_released() -> Result<(), Box<dyn std::error::Error>> {
let config = HostConfig::new("127.0.0.1".parse()?, 4433, Vec::new(), Vec::new());
let engine = HandshakeEngine::new(Registry::builtin(), Arc::new(config));
let first = engine.try_reserve_guest_id(1).await?.unwrap();
assert!(engine.try_reserve_guest_id(1).await?.is_none());
drop(first);
assert!(engine.try_reserve_guest_id(1).await?.is_some());
Ok(())
}
#[tokio::test]
async fn zero_is_rejected_as_a_guest_id() -> Result<(), Box<dyn std::error::Error>> {
let config = HostConfig::new("127.0.0.1".parse()?, 4433, Vec::new(), Vec::new())
.with_guest_id_generator(Box::new(|| Box::pin(async { Some(0) })));
let engine = HandshakeEngine::new(Registry::builtin(), Arc::new(config));
assert!(engine.assign_guest_id().await.is_err());
Ok(())
}
}