Brought Example up to spec
Some checks failed
CI / checks (push) Failing after 3m29s

This commit is contained in:
Alex Emmet 2026-07-19 00:29:24 +02:00
commit 1b796d0ce7
46 changed files with 1755 additions and 691 deletions

View file

@ -1,4 +1,4 @@
#[cfg(any(feature = "crypto", feature = "pipes"))]
#[cfg(feature = "crypto")]
use mtp_codec::{CommunicationType, CommunicationValue};
use mtp_codec::{
DataType, DataValue, Version,
@ -7,6 +7,7 @@ use mtp_codec::{
use mtp_common::RejectionReason;
use mtp_transport::{Receiver, Sender};
use std::sync::Arc;
use std::time::Instant;
#[cfg(feature = "pipes")]
use tokio::sync::mpsc;
@ -18,9 +19,9 @@ use crate::connection::MTPConnection;
use crate::error::AuthState;
use crate::error::{AcceptError, extract_version, send_accepted, send_rejection};
#[cfg(feature = "pipes")]
use crate::pipe::run_dispatcher;
use crate::pipe::PipeDispatcher;
#[cfg(feature = "pipes")]
use crate::pipe::{PipeDispatcher, PipeRequest};
use crate::pipe::run_dispatcher;
pub struct MTPHost {
pub(crate) transport: mtp_transport::Host,
@ -71,11 +72,16 @@ impl MTPHost {
}
if self.handshakes.is_empty() {
let incoming_started = Instant::now();
match self.transport.next().await {
Some((sender, receiver)) => {
tracing::debug!(elapsed = ?incoming_started.elapsed(), "host accept loop: dispatch authentication handshake");
let context = self.context.clone();
self.handshakes.spawn(async move {
context.accept_pair_timed(sender, receiver).await
let handshake_started = Instant::now();
let result = context.accept_pair_timed(sender, receiver).await;
tracing::debug!(elapsed = ?handshake_started.elapsed(), success = result.is_ok(), "host accept loop: authentication handshake finished");
result
});
continue;
}
@ -99,9 +105,15 @@ impl MTPHost {
incoming = self.transport.next() => {
match incoming {
Some((sender, receiver)) => {
tracing::debug!("host accept loop: dispatch authentication handshake");
let context = self.context.clone();
self.handshakes
.spawn(async move { context.accept_pair_timed(sender, receiver).await });
.spawn(async move {
let handshake_started = Instant::now();
let result = context.accept_pair_timed(sender, receiver).await;
tracing::debug!(elapsed = ?handshake_started.elapsed(), success = result.is_ok(), "host accept loop: authentication handshake finished");
result
});
}
None => self.transport_closed = true,
}
@ -172,7 +184,7 @@ impl HandshakeContext {
},
)
.await;
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"authentication not allowed on this host".into(),
));
@ -187,7 +199,7 @@ impl HandshakeContext {
},
)
.await;
sender.close();
sender.close().await;
return Err(AcceptError::MissingVersion);
}
};
@ -208,7 +220,7 @@ impl HandshakeContext {
},
)
.await;
sender.close();
sender.close().await;
return Err(AcceptError::UnsupportedVersion(client_version));
}
};
@ -256,7 +268,7 @@ impl HandshakeContext {
},
)
.await;
sender.close();
sender.close().await;
return Err(AcceptError::MissingVersion);
}
};
@ -273,7 +285,7 @@ impl HandshakeContext {
},
)
.await;
sender.close();
sender.close().await;
return Err(AcceptError::UnsupportedVersion(client_version));
}
};
@ -359,14 +371,6 @@ impl HandshakeContext {
receiver.respond_to_pings(sender.clone());
}
#[cfg(feature = "pipes")]
let (_, app_rx) = mpsc::channel::<
Result<mtp_codec::CommunicationValue, mtp_common::CommunicationError>,
>(1);
#[cfg(feature = "pipes")]
let (_, pipe_req_rx) = mpsc::channel::<PipeRequest>(1);
#[cfg(feature = "pipes")]
let dispatcher = Arc::new(PipeDispatcher);
let task = tokio::spawn(async {});
MTPConnection {
@ -375,12 +379,7 @@ impl HandshakeContext {
sender,
receiver,
path: "/".to_string(),
#[cfg(feature = "pipes")]
app_rx: tokio::sync::Mutex::new(app_rx),
#[cfg(feature = "pipes")]
pipe_req_rx: tokio::sync::Mutex::new(pipe_req_rx),
#[cfg(feature = "pipes")]
pipe_dispatcher: dispatcher,
_pipe_stream: std::marker::PhantomData,
description,
_dispatcher_task: task,
}
@ -451,11 +450,6 @@ impl HandshakeContext {
receiver.respond_to_pings(sender.clone());
}
let (_, app_rx) = mpsc::channel::<
Result<mtp_codec::CommunicationValue, mtp_common::CommunicationError>,
>(1);
let (_, pipe_req_rx) = mpsc::channel::<PipeRequest>(1);
let dispatcher = Arc::new(PipeDispatcher);
let task = tokio::spawn(async {});
MTPConnection {
@ -464,9 +458,7 @@ impl HandshakeContext {
sender,
receiver,
path: "/".to_string(),
app_rx: tokio::sync::Mutex::new(app_rx),
pipe_req_rx: tokio::sync::Mutex::new(pipe_req_rx),
pipe_dispatcher: dispatcher,
_pipe_stream: std::marker::PhantomData,
description,
_dispatcher_task: task,
auth_state,
@ -477,60 +469,6 @@ impl HandshakeContext {
}
}
#[cfg(feature = "pipes")]
impl MTPConnection {
pub async fn receive(
&self,
) -> Result<mtp_codec::CommunicationValue, mtp_common::CommunicationError> {
let mut rx = self.app_rx.lock().await;
match rx.recv().await {
Some(Ok(mut message)) => {
message.set_type_map(self.codec.type_map());
Ok(message)
}
Some(Err(error)) => Err(error),
None => Err(mtp_common::CommunicationError::StreamClosed),
}
}
pub async fn create_pipe(
&self,
description: &str,
) -> Result<crate::pipe::PipeHandle, mtp_common::PipeError> {
let pipe_id = rand::random::<u32>();
let (tx, rx) = tokio::sync::oneshot::channel();
{
let mut pending = self.pipe_dispatcher.pending_creations.lock().await;
pending.insert(pipe_id, tx);
}
let request = CommunicationValue::new(CommunicationType::PipeRequest)
.with_id(pipe_id)
.add_typed_default(DataType::Description, DataValue::Str(description.into()));
self.sender
.send(&request)
.await
.map_err(mtp_common::PipeError::from)?;
Ok(crate::pipe::PipeHandle {
pipe_id,
description: description.to_string(),
sender: self.sender.clone(),
response_rx: rx,
})
}
pub async fn receive_pipe(&self) -> Result<PipeRequest, mtp_common::CommunicationError> {
let mut rx = self.pipe_req_rx.lock().await;
match rx.recv().await {
Some(req) => Ok(req),
None => Err(mtp_common::CommunicationError::StreamClosed),
}
}
}
#[cfg(feature = "crypto")]
enum Flow {
Login {
@ -581,21 +519,21 @@ impl HandshakeContext {
let hello = match receiver.receive().await {
Ok(m) => m,
Err(e) => {
sender.close();
sender.close().await;
return Err(AcceptError::Receive(e));
}
};
let version_str = match hello.get_data(DataType::Version) {
DataValue::Str(s) => s.clone(),
_ => {
sender.close();
sender.close().await;
return Err(AcceptError::MissingVersion);
}
};
let client_version = match Version::parse(&version_str) {
Some(v) => v,
None => {
sender.close();
sender.close().await;
return Err(AcceptError::MissingVersion);
}
};
@ -611,7 +549,7 @@ impl HandshakeContext {
let cid = match hello.get_data(DataType::Id) {
DataValue::UnsignedNumber(n) => *n as u64,
_ => {
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"missing client id".into(),
));
@ -628,7 +566,7 @@ impl HandshakeContext {
DataValue::Str("unknown client id".into()),
);
let _ = sender.send(&rejection).await;
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"unknown client id".into(),
));
@ -644,7 +582,7 @@ impl HandshakeContext {
AcceptError::AuthenticationFailed("invalid public key bundle".into())
})?,
_ => {
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"missing public keys".into(),
));
@ -656,7 +594,7 @@ impl HandshakeContext {
CommunicationType::RegisterResponse,
)
} else {
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"unexpected authentication message".into(),
));
@ -685,11 +623,11 @@ impl HandshakeContext {
client_version: Version,
description: Option<String>,
) -> Result<Option<MTPConnection>, AcceptError> {
use mtp_crypto::{
Ed25519Signer, MlDsaSigner, SignatureScheme, auth, verify_ed25519, verify_ml_dsa,
};
use mtp_crypto::{Ed25519Signer, MlDsaSigner, SignatureScheme, auth, verify_ed25519};
let handshake_started = Instant::now();
let tm = mtp_codec::TypeMap::latest();
let negotiate_started = Instant::now();
let negotiated = match self
.registry
.negotiate(std::slice::from_ref(&client_version))
@ -703,10 +641,11 @@ impl HandshakeContext {
},
)
.await;
sender.close();
sender.close().await;
return Err(AcceptError::UnsupportedVersion(client_version));
}
};
tracing::debug!(elapsed = ?negotiate_started.elapsed(), "authentication handshake: version negotiation");
let pq_enabled = !self
.config
.host_keyring
@ -729,30 +668,43 @@ impl HandshakeContext {
},
)
.await;
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"PQ authentication is required but the host PQ key is absent".into(),
));
}
let host_sign = |payload: &[u8]| -> Result<(Vec<u8>, Vec<u8>), AcceptError> {
let signer = Ed25519Signer::new(&self.config.host_keyring.sig_cl_secret_key)
.map_err(|e| AcceptError::AuthenticationFailed(e.to_string()))?;
let sig = signer
.sign(payload)
.map_err(|e| AcceptError::AuthenticationFailed(e.to_string()))?;
let pq_sig = if pq_enabled {
let pq = MlDsaSigner::new(
let signer_init_started = Instant::now();
let host_pq_signer = if pq_enabled {
Some(Arc::new(
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()))?,
))
} else {
None
};
tracing::debug!(elapsed = ?signer_init_started.elapsed(), "authentication handshake: signer initialization");
let host_sign = |payload: Vec<u8>| async {
let signer = Ed25519Signer::new(&self.config.host_keyring.sig_cl_secret_key)
.map_err(|e| AcceptError::AuthenticationFailed(e.to_string()))?;
pq.sign(payload)
.map_err(|e| AcceptError::AuthenticationFailed(e.to_string()))?
if let Some(pq_signer) = host_pq_signer.as_ref() {
mtp_crypto::sign_parallel::sign_dual_parallel_shared_pq(
signer,
Arc::clone(pq_signer),
payload,
)
.await
.map_err(|e| AcceptError::AuthenticationFailed(e.to_string()))
} else {
Vec::new()
};
Ok((sig, pq_sig))
let sig = signer
.sign(&payload)
.map_err(|e| AcceptError::AuthenticationFailed(e.to_string()))?;
Ok((sig, Vec::new()))
}
};
let challenge_id = match &flow {
@ -761,8 +713,10 @@ impl HandshakeContext {
};
let server_challenge: u128 = rand::random();
let sign_challenge_started = Instant::now();
let (chal_sig, chal_pq_sig) =
host_sign(&auth::challenge_payload(challenge_id, server_challenge))?;
host_sign(auth::challenge_payload(challenge_id, server_challenge)).await?;
tracing::debug!(elapsed = ?sign_challenge_started.elapsed(), "authentication handshake: sign challenge");
let mut challenge_msg = CommunicationValue::new(CommunicationType::Challenge)
.add_typed_default(
@ -782,20 +736,24 @@ impl HandshakeContext {
challenge_msg = challenge_msg
.add_typed_default(DataType::PqSignature, DataValue::Bytes(chal_pq_sig));
}
let send_challenge_started = Instant::now();
if let Err(e) = sender.send(&challenge_msg).await {
sender.close();
sender.close().await;
return Err(AcceptError::Send(e));
}
tracing::debug!(elapsed = ?send_challenge_started.elapsed(), "authentication handshake: send challenge");
let receive_proof_started = Instant::now();
let proof = match receiver.receive().await {
Ok(m) => m,
Err(e) => {
sender.close();
sender.close().await;
return Err(AcceptError::Receive(e));
}
};
tracing::debug!(elapsed = ?receive_proof_started.elapsed(), "authentication handshake: receive client proof");
if Some(proof.get_type()) != CommunicationType::ChallengeResponse.try_to_id(&tm) {
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"missing challenge response".into(),
));
@ -803,7 +761,7 @@ impl HandshakeContext {
let client_nonce = match proof.get_data(DataType::ClientNonce) {
DataValue::UnsignedNumber(n) => *n,
_ => {
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"missing client nonce".into(),
));
@ -812,7 +770,7 @@ impl HandshakeContext {
let sig_bytes = match proof.get_data(DataType::Signature) {
DataValue::Bytes(b) => b.clone(),
_ => {
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"missing challenge signature".into(),
));
@ -837,18 +795,24 @@ impl HandshakeContext {
};
let has_client_pq_key = !bundle.sig_pq_public_key.as_bytes().is_empty();
let pq_ok = if self.config.require_pq {
has_client_pq_key
&& !pq_sig_bytes.is_empty()
&& verify_ml_dsa(&bundle.sig_pq_public_key, &proof_payload, &pq_sig_bytes).is_ok()
let verify_proof_started = Instant::now();
let proof_ok = if pq_sig_bytes.is_empty() {
!self.config.require_pq
&& verify_ed25519(&bundle.sig_cl_public_key, &proof_payload, &sig_bytes).is_ok()
} else if has_client_pq_key {
mtp_crypto::sign_parallel::verify_dual_parallel(
bundle.sig_cl_public_key.clone(),
bundle.sig_pq_public_key.clone(),
proof_payload,
sig_bytes,
pq_sig_bytes,
)
.await
.is_ok()
} else {
pq_sig_bytes.is_empty()
|| (has_client_pq_key
&& verify_ml_dsa(&bundle.sig_pq_public_key, &proof_payload, &pq_sig_bytes)
.is_ok())
false
};
let proof_ok =
verify_ed25519(&bundle.sig_cl_public_key, &proof_payload, &sig_bytes).is_ok() && pq_ok;
tracing::debug!(elapsed = ?verify_proof_started.elapsed(), "authentication handshake: verify client proof");
if !proof_ok {
send_rejection(
@ -858,12 +822,13 @@ impl HandshakeContext {
},
)
.await;
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"client proof signature invalid".into(),
));
}
let register_started = Instant::now();
let (assigned_id, client_bundle) = match flow {
Flow::Login { id, bundle } => (id, bundle),
Flow::Register { bundle, .. } => {
@ -872,12 +837,16 @@ impl HandshakeContext {
(new_id, bundle)
}
};
tracing::debug!(elapsed = ?register_started.elapsed(), "authentication handshake: registration callback");
let (host_sig, host_pq_sig) = host_sign(&auth::host_final_payload(
let sign_final_started = Instant::now();
let (host_sig, host_pq_sig) = host_sign(auth::host_final_payload(
assigned_id,
client_nonce,
server_challenge,
))?;
))
.await?;
tracing::debug!(elapsed = ?sign_final_started.elapsed(), "authentication handshake: sign final response");
let mut response = CommunicationValue::new(response_type)
.add_typed_default(DataType::Connected, DataValue::BoolTrue)
@ -894,14 +863,17 @@ impl HandshakeContext {
response.add_typed_default(DataType::PqSignature, DataValue::Bytes(host_pq_sig));
}
let send_final_started = Instant::now();
if let Err(e) = sender.send(&response).await {
sender.close();
sender.close().await;
return Err(AcceptError::Send(e));
}
if let Err(e) = sender.finish_stream().await {
sender.close();
sender.close().await;
return Err(AcceptError::Send(e));
}
tracing::debug!(elapsed = ?send_final_started.elapsed(), "authentication handshake: send final response");
tracing::debug!(elapsed = ?handshake_started.elapsed(), "authentication handshake: complete");
let codec = match VersionedCodec::for_version(self.registry.clone(), negotiated.clone()) {
Some(codec) => codec,
@ -939,14 +911,14 @@ impl HandshakeContext {
let version_str = match hello.get_data(DataType::Version) {
DataValue::Str(s) => s.clone(),
_ => {
sender.close();
sender.close().await;
return Err(AcceptError::MissingVersion);
}
};
let client_version = match Version::parse(&version_str) {
Some(v) => v,
None => {
sender.close();
sender.close().await;
return Err(AcceptError::MissingVersion);
}
};
@ -959,16 +931,16 @@ impl HandshakeContext {
if Some(hello.get_type()) == CommunicationType::Register.try_to_id(&tm) {
let bundle = match hello.get_data(DataType::PublicKeys) {
DataValue::Bytes(b) => PublicKeyBundle::from_bytes(b).map_err(|_| {
sender.close();
AcceptError::AuthenticationFailed("invalid public key bundle".into())
})?,
_ => {
sender.close();
sender.close().await;
return Err(AcceptError::AuthenticationFailed(
"missing public keys".into(),
));
}
};
sender.close().await;
let pk_bytes = bundle.as_bytes();
return self
.complete_auth_handshake(
@ -1019,7 +991,7 @@ impl HandshakeContext {
},
)
.await;
sender.close();
sender.close().await;
return Err(AcceptError::UnsupportedVersion(client_version));
}
};
@ -1046,7 +1018,7 @@ impl HandshakeContext {
)));
}
sender.close();
sender.close().await;
Err(AcceptError::AuthenticationFailed(
"unexpected message type".into(),
))