diff --git a/Cargo.lock b/Cargo.lock index edf7cb46..8bbd5862 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2318,7 +2318,7 @@ dependencies = [ "prost-reflect-build", "prost-wkt-types", "quinn", - "quinn-plaintext", + "quinn-proto", "quote", "rand 0.8.5", "rcgen", @@ -2330,6 +2330,7 @@ dependencies = [ "rstest", "rust-i18n", "rustls", + "seahash", "serde", "serde_json", "serial_test", @@ -7022,18 +7023,6 @@ dependencies = [ "web-time", ] -[[package]] -name = "quinn-plaintext" -version = "0.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f3e617feaeb6493018fa35fc47ae8b630ac8903d8159e9e747018841b99bad3d" -dependencies = [ - "bytes", - "quinn-proto", - "seahash", - "tracing", -] - [[package]] name = "quinn-proto" version = "0.11.12" diff --git a/easytier/Cargo.toml b/easytier/Cargo.toml index a2fdb63b..2df2744b 100644 --- a/easytier/Cargo.toml +++ b/easytier/Cargo.toml @@ -82,7 +82,8 @@ pin-project-lite = "0.2.13" atomic_refcell = "0.1.13" quinn = { version = "0.11.8", optional = true, features = ["ring"] } -quinn-plaintext = { version = "0.3.0", optional = true } +quinn-proto = { version = "0.11.12", optional = true } +seahash = { version = "4.1.0", optional = true } rustls = { version = "0.23.0", features = [ "ring", "tls12" @@ -373,7 +374,7 @@ full = [ "zstd", ] wireguard = ["dep:boringtun", "dep:ring"] -quic = ["dep:quinn", "dep:quinn-plaintext", "dep:rustls", "dep:rcgen"] +quic = ["dep:quinn", "dep:quinn-proto", "dep:seahash", "dep:rustls", "dep:rcgen"] kcp = ["dep:kcp-sys"] mimalloc = ["dep:mimalloc"] aes-gcm = ["dep:aes-gcm"] diff --git a/easytier/src/tunnel/quic.rs b/easytier/src/tunnel/quic.rs index a9ad21e0..da68c66c 100644 --- a/easytier/src/tunnel/quic.rs +++ b/easytier/src/tunnel/quic.rs @@ -23,6 +23,250 @@ use std::{net::SocketAddr, sync::Arc, time::Duration}; use tokio::net::UdpSocket; // region config +mod crypto { + use crate::utils::BoxExt; + use bytes::{Buf, BytesMut}; + use quinn_proto::crypto::{ + ClientConfig, ExportKeyingMaterialError, KeyPair, Keys, ServerConfig, Session, + UnsupportedVersion, + }; + use quinn_proto::transport_parameters::TransportParameters; + use quinn_proto::{ + ConnectError, ConnectionId, Side, TransportError, + crypto::{CryptoError, HeaderKey, PacketKey}, + }; + use seahash::SeaHasher; + use std::any::Any; + use std::{hash::Hasher, sync::Arc}; + use tracing::{error, instrument, trace}; + + #[derive(Debug, Clone, Copy)] + struct CryptoKey; + + impl CryptoKey { + fn header(self) -> KeyPair> { + KeyPair { + local: Box::new(self), + remote: Box::new(self), + } + } + + fn packet(self) -> KeyPair> { + KeyPair { + local: Box::new(self), + remote: Box::new(self), + } + } + + fn keys(self) -> Keys { + Keys { + header: self.header(), + packet: self.packet(), + } + } + } + + impl HeaderKey for CryptoKey { + fn decrypt(&self, _: usize, _: &mut [u8]) {} + fn encrypt(&self, _: usize, _: &mut [u8]) {} + fn sample_size(&self) -> usize { + 0 + } + } + + impl CryptoKey { + fn checksum(slices: &[&[u8]]) -> u64 { + let mut hasher = SeaHasher::default(); + for slice in slices { + hasher.write(&(slice.len() as u64).to_le_bytes()); + hasher.write(slice); + } + hasher.finish() + } + } + + impl PacketKey for CryptoKey { + #[instrument(level = "trace")] + fn encrypt(&self, packet: u64, buf: &mut [u8], header_len: usize) { + let (header, rest) = buf.split_at_mut(header_len); + let (payload, tag) = rest.split_at_mut(rest.len() - self.tag_len()); + let checksum = Self::checksum(&[header, payload]); + tag.copy_from_slice(&checksum.to_be_bytes()); + trace!(checksum, ?header, ?payload, ?tag); + } + + #[instrument(level = "trace")] + fn decrypt( + &self, + packet: u64, + header: &[u8], + payload: &mut BytesMut, + ) -> Result<(), CryptoError> { + let tag = payload.split_off(payload.len() - self.tag_len()).get_u64(); + trace!(tag, ?payload); + let checksum = Self::checksum(&[header, payload]); + if checksum != tag { + error!(tag, checksum, "checksum mismatch"); + return Err(CryptoError); + } + Ok(()) + } + + fn tag_len(&self) -> usize { + 8 + } + + fn confidentiality_limit(&self) -> u64 { + u64::MAX + } + + fn integrity_limit(&self) -> u64 { + 1 << 36 + } + } + + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + enum HandshakeState { + EmitInitial, + EmitHandshake, + Done, + } + + #[derive(Debug)] + struct QuicSession { + side: Side, + state: HandshakeState, + local: TransportParameters, + remote: Option, + } + + impl QuicSession { + fn new(side: Side, params: TransportParameters) -> Self { + Self { + side, + state: HandshakeState::EmitInitial, + local: params, + remote: None, + } + } + } + + impl Session for QuicSession { + fn initial_keys(&self, _: &ConnectionId, _: Side) -> Keys { + CryptoKey.keys() + } + + fn handshake_data(&self) -> Option> { + self.remote.map(|params| params.boxed() as _) + } + + fn peer_identity(&self) -> Option> { + None + } + + fn early_crypto(&self) -> Option<(Box, Box)> { + None + } + + fn early_data_accepted(&self) -> Option { + Some(false) + } + + #[instrument(level = "trace")] + fn is_handshaking(&self) -> bool { + self.remote.is_none() || self.state != HandshakeState::Done + } + + #[instrument(level = "trace")] + fn read_handshake(&mut self, mut buf: &[u8]) -> Result { + if self.remote.is_none() { + self.remote = Some( + TransportParameters::read(self.side, &mut buf) + .expect("failed to read transport parameters"), + ); + } + Ok(true) + } + + #[instrument(level = "trace")] + fn transport_parameters(&self) -> Result, TransportError> { + Ok(self.remote) + } + + #[instrument(level = "trace")] + fn write_handshake(&mut self, buf: &mut Vec) -> Option { + match self.state { + HandshakeState::EmitInitial => { + if self.side.is_client() { + self.local.write(buf); + } + self.state = HandshakeState::EmitHandshake; + Some(CryptoKey.keys()) + } + HandshakeState::EmitHandshake => { + if self.side.is_server() { + self.local.write(buf); + } + self.state = HandshakeState::Done; + Some(CryptoKey.keys()) + } + HandshakeState::Done => None, + } + } + + fn next_1rtt_keys(&mut self) -> Option>> { + Some(CryptoKey.packet()) + } + + fn is_valid_retry(&self, _: &ConnectionId, _: &[u8], _: &[u8]) -> bool { + true + } + + fn export_keying_material( + &self, + _: &mut [u8], + _: &[u8], + _: &[u8], + ) -> Result<(), ExportKeyingMaterialError> { + Ok(()) + } + } + + #[derive(Debug)] + pub struct CryptoConfig; + + impl ClientConfig for CryptoConfig { + #[instrument(level = "trace")] + fn start_session( + self: Arc, + version: u32, + server_name: &str, + params: &TransportParameters, + ) -> Result, ConnectError> { + Ok(Box::new(QuicSession::new(Side::Client, *params))) + } + } + + impl ServerConfig for CryptoConfig { + fn initial_keys(&self, _: u32, _: &ConnectionId) -> Result { + Ok(CryptoKey.keys()) + } + + fn retry_tag(&self, _: u32, _: &ConnectionId, _: &[u8]) -> [u8; 16] { + [0u8; 16] + } + + #[instrument(level = "trace")] + fn start_session( + self: Arc, + version: u32, + params: &TransportParameters, + ) -> Box { + Box::new(QuicSession::new(Side::Server, *params)) + } + } +} + pub fn transport_config() -> Arc { let mut config = TransportConfig::default(); @@ -39,13 +283,13 @@ pub fn transport_config() -> Arc { } pub fn server_config() -> ServerConfig { - let mut config = quinn_plaintext::server_config(); + let mut config = ServerConfig::with_crypto(Arc::new(crypto::CryptoConfig)); config.transport_config(transport_config()); config } pub fn client_config() -> ClientConfig { - let mut config = quinn_plaintext::client_config(); + let mut config = ClientConfig::new(Arc::new(crypto::CryptoConfig)); config.transport_config(transport_config()); config }