Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
86 changes: 58 additions & 28 deletions src/link_transport_impl.rs
Original file line number Diff line number Diff line change
Expand Up @@ -650,35 +650,20 @@ impl LinkTransport for P2pLinkTransport {
endpoint,
|endpoint| async move {
// Wait for an incoming connection
if let Some(peer_conn) = endpoint.accept().await {
if let Some((peer_conn, conn)) = endpoint.accept_with_connection().await {
// Extract SocketAddr from TransportAddr
let socket_addr = peer_conn
.remote_addr
.as_socket_addr()
.unwrap_or_else(|| SocketAddr::from(([0, 0, 0, 0], 0)));

// Get the underlying QUIC connection by address
match endpoint.get_quic_connection(&socket_addr).await {
Ok(Some(conn)) => {
// Extract peer public key from TLS identity
let public_key = conn
.peer_identity()
.and_then(|id| id.downcast::<Vec<u8>>().ok())
.map(|boxed| *boxed);
let link_conn = P2pLinkConn::new(conn, public_key, socket_addr);
Some((Ok(link_conn), endpoint))
}
Ok(None) => {
// Connection not found, try again
Some((
Err(LinkError::ConnectionFailed(
"Connection not found".to_string(),
)),
endpoint,
))
}
Err(e) => Some((Err(LinkError::ConnectionFailed(e.to_string())), endpoint)),
}
// rustls exposes the authenticated identity as a
// certificate vector whose first entry is the RFC 7250
// ML-DSA SPKI, not as a bare `Vec<u8>`.
let public_key =
crate::p2p_endpoint::extract_public_key_bytes_from_connection(&conn);
let link_conn = P2pLinkConn::new(conn, public_key, socket_addr);
Some((Ok(link_conn), endpoint))
} else {
// Endpoint is shutting down
None
Expand Down Expand Up @@ -725,11 +710,8 @@ impl LinkTransport for P2pLinkTransport {
.map_err(|e| LinkError::ConnectionFailed(e.to_string()))?
.ok_or_else(|| LinkError::ConnectionFailed("Connection not found".to_string()))?;

// Extract peer public key from TLS identity
let public_key = conn
.peer_identity()
.and_then(|id| id.downcast::<Vec<u8>>().ok())
.map(|boxed| *boxed);
// Preserve the identity authenticated by this exact connection.
let public_key = crate::p2p_endpoint::extract_public_key_bytes_from_connection(&conn);

Ok(P2pLinkConn::new(conn, public_key, connected_addr))
})
Expand Down Expand Up @@ -1441,6 +1423,54 @@ mod tests {
assert!(state.capabilities.is_empty());
}

#[tokio::test]
async fn accept_uses_the_authoritative_connection_handle() {
let bind_addr: SocketAddr = "127.0.0.1:0".parse().expect("valid bind address");
let server_endpoint = Arc::new(
P2pEndpoint::new(
P2pConfig::builder()
.bind_addr(bind_addr)
.build()
.expect("valid server config"),
)
.await
.expect("server endpoint"),
);
let server_addr = server_endpoint.local_addr().expect("server address");
let server = P2pLinkTransport::from_endpoint(Arc::clone(&server_endpoint));
let client = P2pEndpoint::new(
P2pConfig::builder()
.bind_addr(bind_addr)
.build()
.expect("valid client config"),
)
.await
.expect("client endpoint");

let mut incoming = server.accept(ProtocolId::DEFAULT);
client
.connect(server_addr)
.await
.expect("client connection");
let accepted = tokio::time::timeout(std::time::Duration::from_secs(10), incoming.next())
.await
.expect("accept timed out")
.expect("accept stream ended")
.expect("accepted connection");

assert_eq!(
accepted.remote_addr().ip(),
client.local_addr().unwrap().ip()
);
assert!(
accepted.peer_public_key().is_some(),
"accepted handle must retain its authenticated identity"
);

client.shutdown().await;
server_endpoint.shutdown().await;
}

// =========================================================================
// Phase 3: SharedTransport Tests
// =========================================================================
Expand Down
7 changes: 6 additions & 1 deletion src/masque/connect.rs
Original file line number Diff line number Diff line change
Expand Up @@ -353,7 +353,7 @@ impl ConnectUdpResponse {

/// Encode the response as wire format
///
/// Format: [status (2)] [flags (1)] [addr if present]
/// Format: [status (2)] [flags (1)] [addr] [reason]
pub fn encode(&self) -> Bytes {
let mut buf = BytesMut::new();

Expand Down Expand Up @@ -407,6 +407,11 @@ impl ConnectUdpResponse {
let flags = buf.get_u8();
let has_addr = (flags & 0x01) != 0;
let has_reason = (flags & 0x02) != 0;
if flags & !0x03 != 0 {
return Err(ConnectError::InvalidResponse(
"unsupported response flags".into(),
));
}

let proxy_public_address = if has_addr {
if buf.remaining() < 1 {
Expand Down
1 change: 1 addition & 0 deletions src/masque/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -127,4 +127,5 @@ pub use relay_server::{
pub use relay_session::{
RelayPeerId, RelaySession, RelaySessionConfig, RelaySessionState, RelaySessionStats,
};
pub(crate) use relay_socket::RelayTunnelControl;
pub use relay_socket::{MasqueRelaySocket, RawRelayStreams};
119 changes: 80 additions & 39 deletions src/masque/relay_server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@
//! ```

use bytes::Bytes;
use parking_lot::RwLock as ParkingRwLock;
use parking_lot::{Mutex as ParkingMutex, RwLock as ParkingRwLock};
use std::collections::HashMap;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::Arc;
Expand Down Expand Up @@ -474,6 +474,20 @@ struct Reservation {
released_at: Instant,
}

/// RAII lease keeping a relay's parent QUIC connection out of ordinary
/// peer-table pruning for the lifetime of one CONNECT-UDP forwarding stream.
pub(crate) struct RelayControlConnectionGuard {
server: Arc<MasqueRelayServer>,
stable_id: usize,
}

impl Drop for RelayControlConnectionGuard {
fn drop(&mut self) {
self.server
.unprotect_relay_control_connection(self.stable_id);
}
}

/// MASQUE Relay Server
///
/// Manages multiple relay sessions and coordinates datagram forwarding
Expand Down Expand Up @@ -526,6 +540,14 @@ pub struct MasqueRelayServer {
/// aborting its reader/writer tasks, so the socket is exclusively owned by the
/// time it can be leased.
forwarding: RwLock<HashMap<u64, ForwardingControl>>,
/// QUIC connections currently carrying one or more CONNECT-UDP forwarding
/// streams, keyed by Quinn's stable connection id.
///
/// Ordinary peer-table maintenance must not force-close these connections:
/// doing so destroys an otherwise healthy, canary-verified relay. Values are
/// reference counts because one authenticated connection may carry multiple
/// relay streams.
relay_control_connections: ParkingMutex<HashMap<usize, usize>>,
/// Mutex stripes serializing concurrent CONNECTs from the same authenticated
/// identity (ADR-011), indexed by the first fingerprint byte. Length is
/// `RELAY_PEER_LOCK_STRIPES`.
Expand Down Expand Up @@ -589,6 +611,7 @@ impl MasqueRelayServer {
upnp_mappings: RwLock::new(HashMap::new()),
reservations: RwLock::new(HashMap::new()),
forwarding: RwLock::new(HashMap::new()),
relay_control_connections: ParkingMutex::new(HashMap::new()),
peer_locks: (0..RELAY_PEER_LOCK_STRIPES)
.map(|_| Mutex::new(()))
.collect(),
Expand Down Expand Up @@ -661,6 +684,7 @@ impl MasqueRelayServer {
upnp_mappings: RwLock::new(HashMap::new()),
reservations: RwLock::new(HashMap::new()),
forwarding: RwLock::new(HashMap::new()),
relay_control_connections: ParkingMutex::new(HashMap::new()),
peer_locks: (0..RELAY_PEER_LOCK_STRIPES)
.map(|_| Mutex::new(()))
.collect(),
Expand All @@ -682,6 +706,40 @@ impl MasqueRelayServer {
}
}

/// Protect a QUIC connection from ordinary peer-table pruning while it
/// carries a live CONNECT-UDP forwarding stream.
pub(crate) fn protect_relay_control_connection(
self: &Arc<Self>,
stable_id: usize,
) -> RelayControlConnectionGuard {
let mut protected = self.relay_control_connections.lock();
*protected.entry(stable_id).or_insert(0) += 1;
drop(protected);
RelayControlConnectionGuard {
server: Arc::clone(self),
stable_id,
}
}

/// Whether `stable_id` currently carries a live CONNECT-UDP stream.
pub(crate) fn is_relay_control_connection(&self, stable_id: usize) -> bool {
self.relay_control_connections
.lock()
.get(&stable_id)
.is_some_and(|count| *count > 0)
}

fn unprotect_relay_control_connection(&self, stable_id: usize) {
let mut protected = self.relay_control_connections.lock();
let Some(count) = protected.get_mut(&stable_id) else {
return;
};
*count -= 1;
if *count == 0 {
protected.remove(&stable_id);
}
}

/// Whether this node is willing to serve as a relay.
pub fn is_relay_serving_enabled(&self) -> bool {
self.relay_serving_enabled.load(Ordering::Acquire)
Expand Down Expand Up @@ -1321,12 +1379,6 @@ impl MasqueRelayServer {
match socket.recv_from(&mut buf).await {
Ok((len, source)) => {
let payload = Bytes::copy_from_slice(&buf[..len]);
tracing::trace!(
session_id,
source = %source,
len,
"RELAY_TUNNEL[srv]: dgram-loop dir1 recv UDP → forwarding to relay-client"
);

// Encode as uncompressed datagram (includes source address
// so client can decode without context registration)
Expand Down Expand Up @@ -1394,24 +1446,10 @@ impl MasqueRelayServer {
};
match resolved {
Some((target, payload)) => {
tracing::trace!(
session_id,
target = %target,
len = payload.len(),
"RELAY_TUNNEL[srv]: dgram-loop dir2 recv from relay-client → sendto target"
);
server2.stats.record_bytes(payload.len() as u64);
server2.stats.record_datagram();
match socket2.send_to(&payload, target).await {
Ok(n) => {
tracing::trace!(
session_id,
target = %target,
len = payload.len(),
sent = n,
"RELAY_TUNNEL[srv]: dgram-loop dir2 sendto OK"
);
}
Ok(_) => {}
Err(e) => {
tracing::warn!(
session_id,
Expand Down Expand Up @@ -1553,10 +1591,6 @@ impl MasqueRelayServer {
match socket.recv_from(&mut buf).await {
Ok((len, source)) => {
let payload = Bytes::copy_from_slice(&buf[..len]);
tracing::trace!(
session_id, source = %source, len,
"RELAY_TUNNEL[srv]: stream-loop dir1 recv UDP → forwarding to relay-client"
);
let datagram =
UncompressedDatagram::new(VarInt::from_u32(0), source, payload);
let encoded = datagram.encode();
Expand Down Expand Up @@ -1717,26 +1751,14 @@ impl MasqueRelayServer {
let mut cursor = Bytes::from(frame_buf);
match UncompressedDatagram::decode(&mut cursor) {
Ok(datagram) => {
tracing::trace!(
session_id, target = %datagram.target,
len = datagram.payload.len(),
"RELAY_TUNNEL[srv]: stream-loop dir2 recv from relay-client → sendto target"
);
stats2.record_bytes(datagram.payload.len() as u64);
stats2.record_datagram();
let target = datagram.target;
let payload_len = datagram.payload.len();
match socket2.send_to(&datagram.payload, target).await {
Ok(n) => {
Ok(_) => {
// Confirmed forwarded to the third-party target.
stats2.record_forwarded_to_target(payload_len as u64, 1);
tracing::trace!(
session_id,
target = %target,
len = payload_len,
sent = n,
"RELAY_TUNNEL[srv]: stream-loop dir2 sendto OK"
);
}
Err(e) if is_message_too_large(&e) => {
// Path-MTU exceeded. Emit a PmtuUpdate
Expand Down Expand Up @@ -2225,6 +2247,25 @@ pub struct SessionInfo {
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;

#[test]
fn relay_control_connection_protection_is_reference_counted() {
let server = Arc::new(MasqueRelayServer::new(
MasqueRelayConfig::default(),
test_addr(9000),
));
let stable_id = 42;

assert!(!server.is_relay_control_connection(stable_id));
let first = server.protect_relay_control_connection(stable_id);
let second = server.protect_relay_control_connection(stable_id);
assert!(server.is_relay_control_connection(stable_id));

drop(first);
assert!(server.is_relay_control_connection(stable_id));
drop(second);
assert!(!server.is_relay_control_connection(stable_id));
}
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};

fn test_addr(port: u16) -> SocketAddr {
Expand Down
Loading
Loading