use crate::network::config::{NetworkConfig, ServerIdentity};
use crate::network::e2e;
use crate::network::protocol::Message;
use anyhow::{Context, Result};
use async_trait::async_trait;
pub use quinn::Connection;
use quinn::{ClientConfig, Endpoint, RecvStream, SendStream, ServerConfig};
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::io::AsyncWriteExt;
#[async_trait]
pub trait MessageSink: Send {
async fn send(&mut self, msg: &Message) -> Result<()>;
async fn send_encoded(&mut self, bytes: &[u8]) -> Result<()>;
}
#[async_trait]
pub trait MessageSource: Send {
async fn recv(&mut self) -> Result<Message>;
async fn recv_raw(&mut self) -> Result<Vec<u8>> {
let msg = self.recv().await?;
Ok(msg.encode()?)
}
}
#[async_trait]
pub trait MessageTransport: Send {
async fn send(&mut self, msg: &Message) -> Result<()>;
async fn send_encoded(&mut self, bytes: &[u8]) -> Result<()>;
async fn recv(&mut self) -> Result<Message>;
fn split(self: Box<Self>) -> (Box<dyn MessageSink>, Box<dyn MessageSource>);
}
pub struct SealedSink {
inner: Box<dyn MessageSink>,
direction: e2e::Direction,
}
impl std::fmt::Debug for SealedSink {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SealedSink")
.field("direction", &self.direction)
.finish_non_exhaustive()
}
}
impl SealedSink {
pub fn new(inner: Box<dyn MessageSink>, direction: e2e::Direction) -> Self {
Self { inner, direction }
}
}
impl SealedSource {
pub fn new(inner: Box<dyn MessageSource>, direction: e2e::Direction) -> Self {
Self { inner, direction }
}
}
#[async_trait]
impl MessageSink for SealedSink {
async fn send(&mut self, msg: &Message) -> Result<()> {
let encoded = msg.encode()?;
MessageSink::send_encoded(self, &encoded).await
}
async fn send_encoded(&mut self, bytes: &[u8]) -> Result<()> {
let sealed = self.direction.seal(bytes)?;
self.inner.send_encoded(&sealed).await
}
}
pub struct SealedSource {
inner: Box<dyn MessageSource>,
direction: e2e::Direction,
}
impl std::fmt::Debug for SealedSource {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SealedSource")
.field("direction", &self.direction)
.finish()
}
}
#[async_trait]
impl MessageSource for SealedSource {
async fn recv(&mut self) -> Result<Message> {
let frame = self.inner.recv_raw().await?;
let plain = self.direction.open(&frame)?;
Ok(Message::decode(&plain)?)
}
}
pub async fn write_encoded<W: tokio::io::AsyncWrite + Unpin>(
writer: &mut W,
bytes: &[u8],
) -> Result<()> {
let len = u32::try_from(bytes.len())
.map_err(|_| anyhow::anyhow!("Encoded message does not fit a u32 length"))?;
writer.write_all(&len.to_le_bytes()).await?;
writer.write_all(bytes).await?;
writer.flush().await?;
Ok(())
}
struct QuicSink {
send: SendStream,
_connection: Connection,
}
#[async_trait]
impl MessageSink for QuicSink {
async fn send(&mut self, msg: &Message) -> Result<()> {
self.send_encoded(&msg.encode()?).await
}
async fn send_encoded(&mut self, bytes: &[u8]) -> Result<()> {
write_encoded(&mut self.send, bytes).await
}
}
struct QuicSource {
recv: RecvStream,
_connection: Connection,
}
#[async_trait]
impl MessageSource for QuicSource {
async fn recv(&mut self) -> Result<Message> {
Message::read_framed(&mut self.recv).await
}
async fn recv_raw(&mut self) -> Result<Vec<u8>> {
Message::read_envelope(&mut self.recv).await
}
}
pub struct QuicTransport {
send: SendStream,
recv: RecvStream,
connection: Connection,
}
impl QuicTransport {
pub fn connection(&self) -> Connection {
self.connection.clone()
}
pub fn new(send: SendStream, recv: RecvStream, connection: Connection) -> Self {
Self {
send,
recv,
connection,
}
}
}
#[async_trait]
impl MessageTransport for QuicTransport {
async fn send(&mut self, msg: &Message) -> Result<()> {
write_encoded(&mut self.send, &msg.encode()?).await
}
async fn send_encoded(&mut self, bytes: &[u8]) -> Result<()> {
write_encoded(&mut self.send, bytes).await
}
async fn recv(&mut self) -> Result<Message> {
Message::read_framed(&mut self.recv).await
}
fn split(self: Box<Self>) -> (Box<dyn MessageSink>, Box<dyn MessageSource>) {
(
Box::new(QuicSink {
send: self.send,
_connection: self.connection.clone(),
}),
Box::new(QuicSource {
recv: self.recv,
_connection: self.connection,
}),
)
}
}
pub fn client_endpoint(config: &NetworkConfig, pin: &[u8]) -> Result<Endpoint> {
let mut client_config = ClientConfig::new(Arc::new(
quinn::crypto::rustls::QuicClientConfig::try_from(NetworkConfig::client_tls_config(pin)?)
.map_err(|e| anyhow::anyhow!("The pinned TLS config is not usable by QUIC: {e}"))?,
));
client_config.transport_config(config.transport_config());
let mut endpoint = Endpoint::client("0.0.0.0:0".parse()?)?;
endpoint.set_default_client_config(client_config);
Ok(endpoint)
}
pub fn server_endpoint(
config: &NetworkConfig,
identity: &ServerIdentity,
bind_addr: SocketAddr,
) -> Result<Endpoint> {
let mut server_config = ServerConfig::with_crypto(Arc::new(
quinn::crypto::rustls::QuicServerConfig::try_from(NetworkConfig::server_crypto_config(
identity,
)?)
.map_err(|e| anyhow::anyhow!("The server TLS config is not usable by QUIC: {e}"))?,
));
server_config.transport_config(config.transport_config());
let endpoint = Endpoint::server(server_config, bind_addr)?;
Ok(endpoint)
}
pub async fn resolve(target: &str) -> Result<SocketAddr> {
let mut addrs: Vec<SocketAddr> = tokio::net::lookup_host(target)
.await
.with_context(|| format!("Could not resolve '{target}'"))?
.collect();
addrs.sort_by_key(|a| u8::from(a.is_ipv6()));
addrs
.into_iter()
.next()
.with_context(|| format!("'{target}' resolved to no addresses"))
}
pub async fn connect_direct(
config: &NetworkConfig,
pin: &[u8],
target: &str,
server_name: &str,
) -> Result<QuicTransport> {
let addr = resolve(target).await?;
let endpoint = client_endpoint(config, pin)?;
let connection = endpoint
.connect(addr, server_name)
.with_context(|| format!("Failed to reach {target} ({addr})"))?
.await
.with_context(|| format!("QUIC handshake with {target} failed"))?;
let (send, recv) = connection
.open_bi()
.await
.context("Failed to open bidirectional stream")?;
Ok(QuicTransport::new(send, recv, connection))
}
#[cfg(test)]
mod tests {
use super::*;
use std::future::Future;
fn block_on<F: Future>(f: F) -> F::Output {
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap()
.block_on(f)
}
#[test]
fn resolve_handles_a_literal_address() {
let addr = block_on(resolve("127.0.0.1:5800")).unwrap();
assert_eq!(addr.to_string(), "127.0.0.1:5800");
}
#[test]
fn an_unresolvable_target_says_so() {
let err = block_on(resolve("pcc.invalid.example:5800"))
.unwrap_err()
.to_string();
assert!(err.contains("pcc.invalid.example"), "unhelpful: {err}");
}
}