use std::net::SocketAddr;
use std::sync::Once;
use std::time::Duration;
use bytes::Bytes;
use futures::{SinkExt, StreamExt};
use nym_bridges::transport::quic::{transport_conn, ClientOptions};
use quinn::{Connection, RecvStream, SendStream};
use tokio_util::codec::{FramedRead, FramedWrite, LengthDelimitedCodec};
use tokio_util::sync::CancellationToken;
use crate::error::{DvpnError, Result};
const LENGTH_DELIMITER_BYTELEN: usize = 2;
const CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
static INSTALL_PROVIDER: Once = Once::new();
#[derive(Clone, Debug)]
pub struct BridgeParams {
pub addresses: Vec<SocketAddr>,
pub sni_host: Option<String>,
pub id_pubkey_base64: String,
}
pub(crate) struct QuicBridgeSender {
framed: FramedWrite<SendStream, LengthDelimitedCodec>,
_conn: Connection,
}
pub(crate) struct QuicBridgeReceiver {
framed: FramedRead<RecvStream, LengthDelimitedCodec>,
_conn: Connection,
}
impl QuicBridgeSender {
pub(crate) async fn send(&mut self, packet: &[u8]) -> Result<()> {
self.framed
.send(Bytes::copy_from_slice(packet))
.await
.map_err(|e| DvpnError::Transport(format!("bridge send: {e}")))
}
}
impl QuicBridgeReceiver {
pub(crate) async fn recv(&mut self) -> Result<Vec<u8>> {
match self.framed.next().await {
Some(Ok(frame)) => Ok(frame.to_vec()),
Some(Err(e)) => Err(DvpnError::Transport(format!("bridge recv: {e}"))),
None => Err(DvpnError::Transport("bridge stream closed".into())),
}
}
}
fn framed_codec() -> LengthDelimitedCodec {
LengthDelimitedCodec::builder()
.length_field_length(LENGTH_DELIMITER_BYTELEN)
.new_codec()
}
pub async fn probe(params: &BridgeParams, cancel: &CancellationToken) -> Result<()> {
let (_send, _recv) = connect(params, cancel).await?;
Ok(())
}
pub(crate) async fn connect(
params: &BridgeParams,
cancel: &CancellationToken,
) -> Result<(QuicBridgeSender, QuicBridgeReceiver)> {
INSTALL_PROVIDER.call_once(|| {
let _ = rustls::crypto::ring::default_provider().install_default();
});
let mut addresses = params.addresses.clone();
addresses.sort_by_key(|a| a.is_ipv6());
let options = ClientOptions {
addresses,
host: params.sni_host.clone(),
id_pubkey: params.id_pubkey_base64.clone(),
};
let conn: Connection = tokio::select! {
_ = cancel.cancelled() => return Err(DvpnError::Cancelled),
r = tokio::time::timeout(CONNECT_TIMEOUT, transport_conn(&options)) => {
r.map_err(|_| DvpnError::Bridge("connect timed out".into()))?
.map_err(|e| DvpnError::Bridge(format!("connect: {e}")))?
}
};
let (send, recv) = tokio::select! {
_ = cancel.cancelled() => return Err(DvpnError::Cancelled),
r = tokio::time::timeout(CONNECT_TIMEOUT, conn.open_bi()) => {
r.map_err(|_| DvpnError::Bridge("open_bi timed out".into()))?
.map_err(|e| DvpnError::Bridge(format!("open_bi: {e}")))?
}
};
Ok((
QuicBridgeSender {
framed: FramedWrite::new(send, framed_codec()),
_conn: conn.clone(),
},
QuicBridgeReceiver {
framed: FramedRead::new(recv, framed_codec()),
_conn: conn,
},
))
}
#[cfg(test)]
mod tests {
use super::*;
use base64::prelude::{Engine as _, BASE64_STANDARD};
use futures::{SinkExt, StreamExt};
use nym_bridges::transport::quic::{create_endpoint, ServerConfig};
fn spawn_mock_bridge() -> (SocketAddr, String, String) {
INSTALL_PROVIDER.call_once(|| {
let _ = rustls::crypto::ring::default_provider().install_default();
});
let secret = [7u8; 32];
let cfg = ServerConfig {
identity_key: Some(BASE64_STANDARD.encode(secret)),
listen: "127.0.0.1:0".parse().unwrap(),
..Default::default()
};
let id_pubkey_base64 = cfg.get_id_pubkey().unwrap();
let sni = bs58::encode(BASE64_STANDARD.decode(&id_pubkey_base64).unwrap()).into_string();
let endpoint = create_endpoint(&cfg).unwrap();
let addr = endpoint.local_addr().unwrap();
tokio::spawn(async move {
let _endpoint = endpoint.clone();
if let Some(incoming) = endpoint.accept().await {
if let Ok(conn) = incoming.await {
if let Ok((send, recv)) = conn.accept_bi().await {
let mut w = FramedWrite::new(send, framed_codec());
let mut r = FramedRead::new(recv, framed_codec());
while let Some(Ok(frame)) = r.next().await {
if w.send(Bytes::from(frame.to_vec())).await.is_err() {
break;
}
}
}
}
}
});
(addr, id_pubkey_base64, sni)
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn bridge_framing_and_pinning_roundtrip() {
let (addr, id_pubkey_base64, sni) = spawn_mock_bridge();
let params = BridgeParams {
addresses: vec![addr],
sni_host: Some(sni),
id_pubkey_base64,
};
let cancel = CancellationToken::new();
let (mut sender, mut receiver) = connect(¶ms, &cancel).await.expect("bridge connect");
let wg: Vec<u8> = (0u16..1200).map(|i| (i % 251) as u8).collect();
sender.send(&wg).await.expect("send");
let echo = receiver.recv().await.expect("recv");
assert_eq!(echo, wg, "framed round-trip mismatch");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn bridge_rejects_wrong_pin() {
let (addr, id_pubkey_base64, sni) = spawn_mock_bridge();
let mut bytes = BASE64_STANDARD.decode(&id_pubkey_base64).unwrap();
bytes[0] ^= 0xFF;
let params = BridgeParams {
addresses: vec![addr],
sni_host: Some(sni),
id_pubkey_base64: BASE64_STANDARD.encode(bytes),
};
let cancel = CancellationToken::new();
assert!(
connect(¶ms, &cancel).await.is_err(),
"connection must be rejected with a mismatched pin"
);
}
}