use std::sync::Arc;
use boring::ssl::SslContextBuilder;
use tokio_quiche::quic::QuicheConnection;
use tokio_quiche::ApplicationOverQuic;
pub use tokio_quiche::ApplicationOverQuic as QUICApplication;
pub use tokio_quiche::quic::QuicheConnection as QUICConnection;
pub use tokio_quiche::quic::HandshakeInfo as QUICHandshake;
pub use tokio_quiche::BoxError as QUICError;
pub use tokio_quiche::QuicResult as QUICOutcome;
pub use tokio_quiche::QuicConnection as QUICGuard;
use crate::models::{ALPN, Role, Version};
use crate::protocol::common::Error;
use crate::tls::{ECHKeys, Identity, Security, TLSConfig};
use crate::helpers::sync;
pub struct Varint;
impl Varint {
pub const MAXIMUM: u64 = (1 << 62) - 1;
pub const MAX_SIZE: usize = 8;
pub fn len(value: u64) -> usize {
match value {
0..=0x3f => 1,
0x40..=0x3fff => 2,
0x4000..=0x3fff_ffff => 4,
_ => 8,
}
}
pub fn encode(out: &mut impl bytes::BufMut, value: u64) {
debug_assert!(value <= Varint::MAXIMUM, "{value} does not fit a variable-length integer");
match Varint::len(value) {
1 => out.put_u8(value as u8),
2 => out.put_slice(&(value as u16 | 0x4000).to_be_bytes()),
4 => out.put_slice(&(value as u32 | 0x8000_0000).to_be_bytes()),
_ => out.put_slice(&(value | 0xc000_0000_0000_0000).to_be_bytes()),
}
}
pub fn decode(input: &[u8]) -> (usize, u64) {
let Some(first) = input.first() else {
return (0, 0);
};
let length = 1 << (first >> 6);
if input.len() < length {
return (0, 0);
}
let mut value = (first & 0x3f) as u64;
for octet in &input[1..length] {
value = value << 8 | *octet as u64;
}
(length, value)
}
pub fn only(payload: &[u8], name: &str) -> Result<u64, Error> {
let (consumed, value) = Varint::decode(payload);
if consumed == 0 || consumed != payload.len() {
return Err(Error::Protocol(format!("{name} payload is not a single variable-length integer")));
}
Ok(value)
}
}
pub struct QUICStreamID;
impl QUICStreamID {
pub const STEP: u64 = 4;
pub fn is_bidi(id: u64) -> bool {
id & 0x2 == 0
}
pub fn is_uni(id: u64) -> bool {
id & 0x2 != 0
}
pub fn client_initiated(id: u64) -> bool {
id & 0x1 == 0
}
pub fn first_bidi(role: Role) -> u64 {
if role.is_client() { 0 } else { 1 }
}
pub fn first_uni(role: Role) -> u64 {
if role.is_client() { 2 } else { 3 }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StreamRead {
Data {
len: usize,
fin: bool,
},
Done,
Reset(u64),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum StreamWrite {
Sent(usize),
Blocked,
Stopped(u64),
}
pub trait QUICTransport {
fn send(&mut self, stream_id: u64, data: &[u8], fin: bool) -> Result<StreamWrite, Error>;
fn receive(&mut self, stream_id: u64, out: &mut [u8]) -> Result<StreamRead, Error>;
fn shutdown_read(&mut self, stream_id: u64, code: u64) -> Result<(), Error>;
fn shutdown_write(&mut self, stream_id: u64, code: u64) -> Result<(), Error>;
fn readable(&self) -> impl Iterator<Item = u64>;
fn close(&mut self, code: u64, reason: &[u8]) -> Result<(), Error>;
fn application_protocol(&self) -> &[u8];
fn version(&self) -> u32;
}
impl QUICTransport for QuicheConnection {
fn send(&mut self, stream_id: u64, data: &[u8], fin: bool) -> Result<StreamWrite, Error> {
match self.stream_send(stream_id, data, fin) {
Ok(sent) => Ok(StreamWrite::Sent(sent)),
Err(quiche::Error::Done) => Ok(StreamWrite::Blocked),
Err(quiche::Error::StreamStopped(code)) => Ok(StreamWrite::Stopped(code)),
Err(err) => Err(Error::quic(err)),
}
}
fn receive(&mut self, stream_id: u64, out: &mut [u8]) -> Result<StreamRead, Error> {
match self.stream_recv(stream_id, out) {
Ok((len, fin)) => Ok(StreamRead::Data { len, fin }),
Err(quiche::Error::Done) => Ok(StreamRead::Done),
Err(quiche::Error::StreamReset(code)) => Ok(StreamRead::Reset(code)),
Err(err) => Err(Error::quic(err)),
}
}
fn shutdown_read(&mut self, stream_id: u64, code: u64) -> Result<(), Error> {
match self.stream_shutdown(stream_id, quiche::Shutdown::Read, code) {
Ok(()) | Err(quiche::Error::Done) => Ok(()),
Err(err) => Err(Error::quic(err)),
}
}
fn shutdown_write(&mut self, stream_id: u64, code: u64) -> Result<(), Error> {
match self.stream_shutdown(stream_id, quiche::Shutdown::Write, code) {
Ok(()) | Err(quiche::Error::Done) => Ok(()),
Err(err) => Err(Error::quic(err)),
}
}
fn readable(&self) -> impl Iterator<Item = u64> {
QuicheConnection::readable(self)
}
fn close(&mut self, code: u64, reason: &[u8]) -> Result<(), Error> {
match QuicheConnection::close(self, true, code, reason) {
Ok(()) | Err(quiche::Error::Done) => Ok(()),
Err(err) => Err(Error::quic(err)),
}
}
fn application_protocol(&self) -> &[u8] {
self.application_proto()
}
fn version(&self) -> u32 {
quiche::PROTOCOL_VERSION
}
}
pub struct Handshake {
pub alpn: Vec<u8>,
pub version: u32,
}
impl Handshake {
pub fn of(transport: &impl QUICTransport) -> Self {
Self { alpn: transport.application_protocol().to_vec(), version: transport.version() }
}
pub fn negotiated(&self, versions: &[Version]) -> Result<Version, Error> {
ALPN::negotiated((!self.alpn.is_empty()).then_some(&self.alpn), versions)
}
pub fn security(&self) -> Security {
Security::quic(Some(self.version))
}
}
pub struct QUICConfig {
pub versions: Vec<Version>,
pub idle_timeout: f64,
pub max_streams_bidi: Option<u64>,
pub enable_dgram: bool,
}
impl QUICConfig {
pub fn settings(&self) -> tokio_quiche::settings::QuicSettings {
let mut settings = tokio_quiche::settings::QuicSettings::default();
settings.alpn = ALPN::list(&self.versions);
settings.max_idle_timeout = sync::Timeout::duration(self.idle_timeout);
settings.enable_dgram = self.enable_dgram;
if let Some(max) = self.max_streams_bidi {
settings.initial_max_streams_bidi = max;
}
settings
}
pub fn placeholder_certificate() -> tokio_quiche::settings::TlsCertificatePaths<'static> {
tokio_quiche::settings::TlsCertificatePaths { cert: "", private_key: "", kind: tokio_quiche::settings::CertificateKind::X509 }
}
}
pub type QUICIncoming = tokio_quiche::InitialQuicConnection<tokio::net::UdpSocket, tokio_quiche::metrics::DefaultMetrics>;
pub type QUICIncomingStream = tokio::sync::mpsc::Receiver<std::io::Result<QUICIncoming>>;
pub struct QUICListener;
impl QUICListener {
pub fn bind(udp: std::net::UdpSocket, config: &QUICConfig, hook: Arc<dyn tokio_quiche::quic::ConnectionHook + Send + Sync>) -> Result<(QUICIncomingStream, std::net::SocketAddr), Error> {
let address = udp.local_addr()?;
let hooks = tokio_quiche::settings::Hooks { connection_hook: Some(hook) };
let params = tokio_quiche::ConnectionParams::new_server(config.settings(), QUICConfig::placeholder_certificate(), hooks);
let listeners = tokio_quiche::listen([udp], params, tokio_quiche::metrics::DefaultMetrics).map_err(Error::IO)?;
let incoming = listeners.into_iter().next().ok_or(Error::Closed)?.into_inner();
Ok((incoming, address))
}
}
pub struct QUICDialer;
impl QUICDialer {
pub async fn connect(host: &str, udp: tokio::net::UdpSocket, config: &QUICConfig, hook: Arc<dyn tokio_quiche::quic::ConnectionHook + Send + Sync>, application: impl ApplicationOverQuic) -> Result<tokio_quiche::QuicConnection, Error> {
let socket: tokio_quiche::socket::Socket<Arc<tokio::net::UdpSocket>, Arc<tokio::net::UdpSocket>> = udp.try_into().map_err(Error::IO)?;
let hooks = tokio_quiche::settings::Hooks { connection_hook: Some(hook) };
let params = tokio_quiche::ConnectionParams::new_client(config.settings(), Some(QUICConfig::placeholder_certificate()), hooks);
tokio_quiche::quic::connect_with_config(socket, Some(host), ¶ms, application)
.await
.map_err(|err| Error::TLS(err.to_string()))
}
}
pub struct QUICServerTLS {
pub identity: Identity,
pub ech: Option<ECHKeys>,
pub tls: TLSConfig,
}
impl tokio_quiche::quic::ConnectionHook for QUICServerTLS {
fn create_custom_ssl_context_builder(&self, _settings: tokio_quiche::settings::TlsCertificatePaths<'_>) -> Option<SslContextBuilder> {
self.tls.quic_server(&self.identity, self.ech.as_ref()).ok()
}
}
pub struct QUICClientTLS {
pub roots: Vec<Vec<u8>>,
pub tls: TLSConfig,
}
impl tokio_quiche::quic::ConnectionHook for QUICClientTLS {
fn create_custom_ssl_context_builder(&self, _settings: tokio_quiche::settings::TlsCertificatePaths<'_>) -> Option<SslContextBuilder> {
self.tls.quic_client(&self.roots).ok()
}
}