1use std::{collections::HashMap, io, net::SocketAddr, sync::Arc, time::Duration};
8
9use crate::connection::Connection;
10use nxtquic_crypto::{
11 CryptoProvider, DirectionalKeys, HandshakeState, HeaderProtectionKey, Keys, PacketKey,
12 RustlsCryptoProvider,
13};
14use nxtquic_proto::{
15 ConnectionId, StreamId, Version,
16 frame::{AckFrame, CryptoFrame, Frame, StreamFrame},
17 packet::{
18 PacketType, decode_header,
19 initial::parse_initial,
20 protection::{apply_header_protection, remove_header_protection},
21 },
22 varint::VarInt,
23};
24
25#[derive(Default, Clone, Debug)]
27pub struct EndpointConfig;
28
29#[derive(Clone)]
31pub struct ServerConfig {
32 rustls: Arc<rustls::ServerConfig>,
33 transport: TransportConfig,
34}
35
36impl std::fmt::Debug for ServerConfig {
37 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
38 formatter
39 .debug_struct("ServerConfig")
40 .field("transport", &self.transport)
41 .finish_non_exhaustive()
42 }
43}
44
45impl ServerConfig {
46 pub fn builder() -> ServerConfigBuilder {
48 ServerConfigBuilder::default()
49 }
50
51 pub fn rustls(&self) -> &Arc<rustls::ServerConfig> {
53 &self.rustls
54 }
55
56 pub fn transport(&self) -> &TransportConfig {
58 &self.transport
59 }
60}
61
62#[derive(Default)]
64pub struct ServerConfigBuilder {
65 rustls: Option<Arc<rustls::ServerConfig>>,
66 transport: TransportConfig,
67}
68
69impl ServerConfigBuilder {
70 pub fn with_rustls(mut self, config: Arc<rustls::ServerConfig>) -> io::Result<Self> {
72 if config
73 .alpn_protocols
74 .iter()
75 .all(|protocol| protocol.as_slice() != b"h3")
76 {
77 return Err(io::Error::new(
78 io::ErrorKind::InvalidInput,
79 "a QUIC HTTP/3 server configuration must advertise ALPN h3",
80 ));
81 }
82 self.rustls = Some(config);
83 Ok(self)
84 }
85
86 pub fn with_transport_config(mut self, configure: impl FnOnce(&mut TransportConfig)) -> Self {
88 configure(&mut self.transport);
89 self
90 }
91
92 pub fn build(self) -> io::Result<ServerConfig> {
94 let rustls = self.rustls.ok_or_else(|| {
95 io::Error::new(
96 io::ErrorKind::InvalidInput,
97 "missing Rustls server configuration",
98 )
99 })?;
100 Ok(ServerConfig {
101 rustls,
102 transport: self.transport,
103 })
104 }
105}
106
107#[derive(Default, Clone, Debug)]
109pub struct TransportConfig {
110 max_idle_timeout: Option<Duration>,
111}
112
113impl TransportConfig {
114 pub fn max_idle_timeout(&mut self, timeout: Option<Duration>) -> &mut Self {
116 self.max_idle_timeout = timeout;
117 self
118 }
119
120 pub fn idle_timeout(&self) -> Option<Duration> {
122 self.max_idle_timeout
123 }
124}
125
126#[derive(Default, Clone, Debug)]
128pub struct ClientConfig;
129
130pub struct Endpoint {
132 config: EndpointConfig,
133 server_config: Option<ServerConfig>,
134 socket: Arc<tokio::net::UdpSocket>,
135 incoming: Arc<tokio::sync::Mutex<tokio::sync::mpsc::Receiver<Incoming>>>,
136}
137
138pub struct Connecting;
140
141pub struct Incoming {
143 remote_addr: SocketAddr,
144 packet: Vec<u8>,
145 initial: DecryptedInitial,
146 socket: Arc<tokio::net::UdpSocket>,
147 server_config: Option<ServerConfig>,
148 datagrams: Option<tokio::sync::mpsc::Receiver<Vec<u8>>>,
149}
150
151#[derive(Debug)]
156struct DecryptedInitial {
157 version: Version,
158 dcid: ConnectionId,
159 scid: ConnectionId,
160 packet_number: u64,
161 frames: Vec<Frame>,
162}
163
164fn decrypt_client_initial(packet: &[u8]) -> io::Result<DecryptedInitial> {
165 let parsed =
166 parse_initial(packet).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
167 let version = match parsed.version {
168 1 => Version::V1,
169 0x6b33_43cf => Version::V2,
170 version => {
171 return Err(io::Error::new(
172 io::ErrorKind::InvalidData,
173 format!("unsupported QUIC version {version:#010x}"),
174 ));
175 }
176 };
177
178 let provider = RustlsCryptoProvider::new(None, None);
179 let keys = provider.initial_keys(version, &parsed.dcid);
180 let mut protected = packet[..parsed.packet_end].to_vec();
181
182 let sample_offset = parsed.packet_number_offset + 4;
183 let sample: [u8; 16] = protected[sample_offset..sample_offset + 16]
184 .try_into()
185 .expect("Initial parser validates the header-protection sample");
186 let mask = keys
187 .client
188 .header_key
189 .protection_mask(&sample)
190 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
191 let packet_number_len =
192 remove_header_protection(&mut protected, &mask, parsed.packet_number_offset);
193 if packet_number_len == 0 || parsed.packet_number_offset + packet_number_len > parsed.packet_end
194 {
195 return Err(io::Error::new(
196 io::ErrorKind::InvalidData,
197 "invalid QUIC Initial packet number",
198 ));
199 }
200
201 let mut packet_number = 0_u64;
202 for byte in
203 &protected[parsed.packet_number_offset..parsed.packet_number_offset + packet_number_len]
204 {
205 packet_number = (packet_number << 8) | u64::from(*byte);
206 }
207 let header_end = parsed.packet_number_offset + packet_number_len;
208 let mut payload = protected[header_end..parsed.packet_end].to_vec();
209 keys.client
210 .packet_key
211 .decrypt(packet_number, &protected[..header_end], &mut payload)
212 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
213
214 let mut frames = Vec::new();
215 let mut payload = &payload[..];
216 while !payload.is_empty() {
217 let before = payload.len();
218 let frame = Frame::decode(&mut payload)
219 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
220 frames.push(frame);
221 if payload.len() == before {
222 return Err(io::Error::new(
223 io::ErrorKind::InvalidData,
224 "QUIC frame decoder made no progress",
225 ));
226 }
227 }
228
229 Ok(DecryptedInitial {
230 version,
231 dcid: parsed.dcid,
232 scid: parsed.scid,
233 packet_number,
234 frames,
235 })
236}
237
238fn encode_server_initial(
239 version: Version,
240 client_scid: ConnectionId,
241 server_scid: ConnectionId,
242 packet_number: u64,
243 crypto_data: Vec<u8>,
244 keys: nxtquic_crypto::InitialKeys,
245) -> io::Result<Vec<u8>> {
246 let mut plaintext = Vec::new();
247 Frame::Crypto(CryptoFrame {
248 offset: VarInt::ZERO,
249 data: crypto_data.into(),
250 })
251 .encode(&mut plaintext);
252
253 let mut header_prefix = Vec::new();
254 header_prefix.push(0xc0); header_prefix.extend_from_slice(&version.as_u32().to_be_bytes());
256 header_prefix.push(client_scid.len() as u8);
257 header_prefix.extend_from_slice(client_scid.as_bytes());
258 header_prefix.push(server_scid.len() as u8);
259 header_prefix.extend_from_slice(server_scid.as_bytes());
260 header_prefix.push(0); loop {
265 let encrypted_len = plaintext.len() + keys.server.packet_key.tag_len();
266 let mut candidate = header_prefix.clone();
267 VarInt::from_u64((1 + encrypted_len) as u64)
268 .expect("QUIC packet lengths fit in a VarInt")
269 .encode(&mut candidate);
270 if candidate.len() + 1 + encrypted_len >= 1200 {
271 header_prefix = candidate;
272 break;
273 }
274 plaintext.push(0);
275 }
276
277 let packet_number_len = 1;
278 let packet_number_offset = header_prefix.len();
279 header_prefix.push(packet_number as u8);
280 keys.server
281 .packet_key
282 .encrypt(packet_number, &header_prefix, &mut plaintext)
283 .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
284 header_prefix.extend_from_slice(&plaintext);
285
286 let sample: [u8; 16] = header_prefix[packet_number_offset + 4..packet_number_offset + 20]
287 .try_into()
288 .expect("1200-byte Initial always contains a header-protection sample");
289 let mask = keys
290 .server
291 .header_key
292 .protection_mask(&sample)
293 .map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
294 apply_header_protection(
295 &mut header_prefix,
296 &mask,
297 packet_number_offset,
298 packet_number_len,
299 );
300 Ok(header_prefix)
301}
302
303impl std::fmt::Debug for Incoming {
304 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
305 formatter
306 .debug_struct("Incoming")
307 .field("remote_addr", &self.remote_addr)
308 .field("packet_len", &self.packet.len())
309 .finish()
310 }
311}
312
313impl Incoming {
314 pub fn remote_address(&self) -> SocketAddr {
316 self.remote_addr
317 }
318
319 pub(crate) fn packet(&self) -> &[u8] {
321 &self.packet
322 }
323}
324
325impl Endpoint {
326 pub async fn bind(addr: SocketAddr) -> io::Result<Self> {
329 Self::bind_inner(EndpointConfig::default(), None, addr).await
330 }
331
332 pub async fn server(server_config: ServerConfig, addr: SocketAddr) -> io::Result<Self> {
334 Self::bind_inner(EndpointConfig::default(), Some(server_config), addr).await
335 }
336
337 async fn bind_inner(
338 config: EndpointConfig,
339 server_config: Option<ServerConfig>,
340 addr: SocketAddr,
341 ) -> io::Result<Self> {
342 let socket = tokio::net::UdpSocket::bind(addr).await?;
343 let socket = Arc::new(socket);
344 let (incoming_tx, incoming_rx) = tokio::sync::mpsc::channel(128);
345 let receive_socket = Arc::clone(&socket);
346 let receive_server_config = server_config.clone();
347 tokio::spawn(async move {
348 let mut packet = vec![0_u8; 65_535];
349 let mut peers: HashMap<SocketAddr, tokio::sync::mpsc::Sender<Vec<u8>>> = HashMap::new();
350 loop {
351 let (length, remote_addr) = match receive_socket.recv_from(&mut packet).await {
352 Ok(received) => received,
353 Err(_) => break,
354 };
355
356 let datagram = &packet[..length];
357 if let Some(peer_tx) = peers.get(&remote_addr) {
358 if peer_tx.send(datagram.to_vec()).await.is_err() {
359 peers.remove(&remote_addr);
360 }
361 continue;
362 }
363 let Ok(header) = decode_header(datagram) else {
364 continue;
365 };
366 if header.packet_type != PacketType::Initial {
367 continue;
368 }
369 let Ok(initial) = decrypt_client_initial(datagram) else {
370 continue;
371 };
372
373 let (peer_tx, peer_rx) = tokio::sync::mpsc::channel(256);
374 peers.insert(remote_addr, peer_tx);
375
376 if incoming_tx
377 .send(Incoming {
378 remote_addr,
379 packet: datagram.to_vec(),
380 initial,
381 socket: Arc::clone(&receive_socket),
382 server_config: receive_server_config.clone(),
383 datagrams: Some(peer_rx),
384 })
385 .await
386 .is_err()
387 {
388 break;
389 }
390 }
391 });
392
393 Ok(Self {
394 config,
395 server_config,
396 socket,
397 incoming: Arc::new(tokio::sync::Mutex::new(incoming_rx)),
398 })
399 }
400
401 pub fn local_addr(&self) -> io::Result<SocketAddr> {
403 self.socket.local_addr()
404 }
405
406 pub fn server_config(&self) -> Option<&ServerConfig> {
408 self.server_config.as_ref()
409 }
410
411 pub async fn connect(
413 &self,
414 _addr: SocketAddr,
415 _server_name: &str,
416 ) -> std::io::Result<Connecting> {
417 Ok(Connecting)
418 }
419
420 pub async fn accept(&self) -> Option<Incoming> {
422 self.incoming.lock().await.recv().await
423 }
424}
425
426impl Connecting {
427 pub async fn await_connection(self) -> std::io::Result<Connection> {
429 Ok(Connection::new())
430 }
431}
432
433impl Incoming {
434 pub async fn accept(self) -> std::io::Result<Connection> {
436 let config = self.server_config.ok_or_else(|| {
437 io::Error::new(
438 io::ErrorKind::InvalidInput,
439 "cannot accept a connection on an endpoint without ServerConfig",
440 )
441 })?;
442 let mut crypto = Vec::new();
443 for frame in &self.initial.frames {
444 if let Frame::Crypto(frame) = frame {
445 if frame.offset.into_inner() != crypto.len() as u64 {
446 return Err(io::Error::new(
447 io::ErrorKind::InvalidData,
448 "Initial CRYPTO frames must be contiguous",
449 ));
450 }
451 crypto.extend_from_slice(&frame.data);
452 }
453 }
454 if crypto.is_empty() {
455 return Err(io::Error::new(
456 io::ErrorKind::InvalidData,
457 "Initial packet contains no CRYPTO frame",
458 ));
459 }
460
461 let provider = RustlsCryptoProvider::new(None, Some(Arc::clone(config.rustls())));
462
463 let mut transport_params = nxtquic_proto::TransportParameters::default();
464 transport_params.initial_max_data = VarInt::from_u64(10 * 1024 * 1024).unwrap();
465 transport_params.initial_max_stream_data_bidi_local = VarInt::from_u64(1024 * 1024).unwrap();
466 transport_params.initial_max_stream_data_bidi_remote = VarInt::from_u64(1024 * 1024).unwrap();
467 transport_params.initial_max_stream_data_uni = VarInt::from_u64(1024 * 1024).unwrap();
468 transport_params.initial_max_streams_bidi = VarInt::from_u64(100).unwrap();
469 transport_params.initial_max_streams_uni = VarInt::from_u64(100).unwrap();
470 if let Some(timeout) = config.transport().idle_timeout() {
471 transport_params.max_idle_timeout = VarInt::from_u64(timeout.as_millis() as u64).unwrap_or(VarInt::from_u32(0));
472 }
473 transport_params.max_udp_payload_size = VarInt::from_u64(1350).unwrap();
474 transport_params.active_connection_id_limit = VarInt::from_u64(2).unwrap();
475
476 let server_scid = ConnectionId::from_slice(&rand::random::<[u8; 16]>());
477 transport_params.initial_source_connection_id = Some(server_scid.clone());
478 transport_params.original_destination_connection_id = Some(self.initial.dcid.clone());
479
480 let mut tp_encoded = Vec::new();
481 transport_params.encode(&mut tp_encoded);
482
483 let mut handshake = provider
484 .new_server_handshake(tp_encoded)
485 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
486 let output = handshake
487 .process(&crypto)
488 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
489 if output.data.is_empty() {
490 return Err(io::Error::new(
491 io::ErrorKind::InvalidData,
492 "TLS produced no Initial handshake data",
493 ));
494 }
495
496 let keys = provider.initial_keys(self.initial.version, &self.initial.dcid);
497 let response = encode_server_initial(
498 self.initial.version,
499 self.initial.scid,
500 server_scid,
501 0,
502 output.data,
503 keys,
504 )?;
505 self.socket.send_to(&response, self.remote_addr).await?;
506 let (outgoing_tx, outgoing_rx) = tokio::sync::mpsc::unbounded_channel();
507 let connection = Connection::with_outgoing(outgoing_tx);
508 if let Some(datagrams) = self.datagrams {
509 let driver_connection = connection.clone();
510 let driver_initial = self.initial;
511 let driver_socket = Arc::clone(&self.socket);
512 let driver_remote = self.remote_addr;
513 let driver_initial_keys = provider.initial_keys(driver_initial.version, &driver_initial.dcid);
514 let mut driver_handshake_keys = None;
515 let mut driver_one_rtt_keys = None;
516 while let Some((level, keys)) = handshake.next_keys() {
517 if level == nxtquic_crypto::Level::OneRtt {
518 driver_one_rtt_keys = Some(keys);
519 } else {
520 driver_handshake_keys = Some((level, keys));
521 }
522 }
523 tokio::spawn(async move {
524 drive_server_connection(
525 driver_connection,
526 driver_socket,
527 driver_remote,
528 datagrams,
529 driver_initial,
530 server_scid,
531 driver_initial_keys,
532 driver_handshake_keys,
533 driver_one_rtt_keys,
534 handshake,
535 outgoing_rx,
536 )
537 .await;
538 });
539 }
540 Ok(connection)
541 }
542}
543
544struct LongPacket {
545 packet_type: PacketType,
546 dcid: ConnectionId,
547 scid: ConnectionId,
548 pn_offset: usize,
549 packet_end: usize,
550}
551
552fn parse_long_packet(datagram: &[u8]) -> io::Result<LongPacket> {
553 if datagram.len() < 7 || datagram[0] & 0xc0 != 0xc0 {
554 return Err(io::Error::new(io::ErrorKind::InvalidData, "invalid QUIC long header"));
555 }
556 let version = u32::from_be_bytes(datagram[1..5].try_into().unwrap());
557 if version == 0 {
558 return Err(io::Error::new(io::ErrorKind::InvalidData, "version negotiation is not a protected packet"));
559 }
560 let packet_type = match (datagram[0] >> 4) & 0x03 {
561 0 => PacketType::Initial,
562 1 => PacketType::ZeroRtt,
563 2 => PacketType::Handshake,
564 _ => PacketType::Retry,
565 };
566 let mut offset = 5;
567 let dcid_len = datagram[offset] as usize;
568 offset += 1;
569 let dcid_end = offset.checked_add(dcid_len).ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid DCID"))?;
570 let dcid = ConnectionId::from_slice(datagram.get(offset..dcid_end).ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "truncated DCID"))?);
571 offset = dcid_end;
572 let scid_len = *datagram.get(offset).ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "missing SCID length"))? as usize;
573 offset += 1;
574 let scid_end = offset.checked_add(scid_len).ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid SCID"))?;
575 let scid = ConnectionId::from_slice(datagram.get(offset..scid_end).ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "truncated SCID"))?);
576 offset = scid_end;
577 if packet_type == PacketType::Initial {
578 let mut rest = &datagram[offset..];
579 let token_len = VarInt::decode(&mut rest).map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid Initial token length"))?.into_inner() as usize;
580 let consumed = datagram[offset..].len() - rest.len();
581 offset += consumed.checked_add(token_len).ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid token"))?;
582 }
583 let mut rest = &datagram[offset..];
584 let payload_len = VarInt::decode(&mut rest).map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "invalid QUIC payload length"))?.into_inner() as usize;
585 offset += datagram[offset..].len() - rest.len();
586 let packet_end = offset.checked_add(payload_len).ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid packet length"))?;
587 if packet_end > datagram.len() || packet_end < offset + 4 + 16 {
588 return Err(io::Error::new(io::ErrorKind::UnexpectedEof, "truncated protected QUIC packet"));
589 }
590 Ok(LongPacket { packet_type, dcid, scid, pn_offset: offset, packet_end })
591}
592
593fn decrypt_payload(mut packet: Vec<u8>, parsed: &LongPacket, keys: &DirectionalKeys) -> io::Result<(u64, Vec<u8>)> {
594 let sample: [u8; 16] = packet[parsed.pn_offset + 4..parsed.pn_offset + 20].try_into().map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "missing header-protection sample"))?;
595 let mask = keys.header_key.protection_mask(&sample).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
596 let pn_len = remove_header_protection(&mut packet, &mask, parsed.pn_offset);
597 if pn_len == 0 || parsed.pn_offset + pn_len > parsed.packet_end {
598 return Err(io::Error::new(io::ErrorKind::InvalidData, "invalid packet number"));
599 }
600 let mut packet_number = 0;
601 for byte in &packet[parsed.pn_offset..parsed.pn_offset + pn_len] {
602 packet_number = (packet_number << 8) | u64::from(*byte);
603 }
604 let header_end = parsed.pn_offset + pn_len;
605 let mut payload = packet[header_end..parsed.packet_end].to_vec();
606 keys.packet_key.decrypt(packet_number, &packet[..header_end], &mut payload).map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
607 Ok((packet_number, payload))
608}
609
610async fn drive_server_connection(
611 connection: Connection,
612 socket: Arc<tokio::net::UdpSocket>,
613 remote_addr: SocketAddr,
614 mut datagrams: tokio::sync::mpsc::Receiver<Vec<u8>>,
615 initial: DecryptedInitial,
616 server_scid: ConnectionId,
617 initial_keys: nxtquic_crypto::InitialKeys,
618 mut handshake_keys: Option<(nxtquic_crypto::Level, Keys)>,
619 mut one_rtt_keys: Option<Keys>,
620 mut handshake: impl HandshakeState,
621 mut outgoing: tokio::sync::mpsc::UnboundedReceiver<crate::stream::WriteCommand>,
622) {
623 let mut largest_handshake = 0;
624 let mut application_packet_number = 0;
625 loop {
626 let Some(event) = (tokio::select! {
627 datagram = datagrams.recv() => datagram.map(Ok),
628 command = outgoing.recv() => command.map(Err),
629 else => None,
630 }) else { break };
631 let datagram = match event {
632 Err(command) => {
633 if let Some(keys) = one_rtt_keys.as_ref() {
634 let crate::stream::WriteCommand::Data { stream_id, offset, data, fin } = command;
635 let mut encoded = Vec::new();
636 Frame::Stream(StreamFrame {
637 stream_id: StreamId::from_u64(stream_id),
638 offset: VarInt::from_u64(offset).unwrap(),
639 length: Some(VarInt::from_u64(data.len() as u64).unwrap()),
640 fin,
641 data: bytes::Bytes::from(data),
642 }).encode(&mut encoded);
643 let _ = send_protected_short(&socket, remote_addr, initial.scid, application_packet_number, &encoded, &keys.local).await;
644 application_packet_number += 1;
645 }
646 continue;
647 }
648 Ok(datagram) => datagram,
649 };
650 let Ok(header) = decode_header(&datagram) else { continue };
651 let (payload, packet_number, level) = match header.packet_type {
652 PacketType::Initial => {
653 let Ok(parsed) = parse_long_packet(&datagram) else { continue };
654 let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &initial_keys.client) else { continue };
655 (payload, pn, nxtquic_crypto::Level::Initial)
656 }
657 PacketType::Handshake => {
658 let Ok(parsed) = parse_long_packet(&datagram) else { continue };
659 let Some((_, keys)) = handshake_keys.as_ref() else { continue };
660 let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &keys.remote) else { continue };
661 largest_handshake = largest_handshake.max(pn);
662 (payload, pn, nxtquic_crypto::Level::Handshake)
663 }
664 PacketType::Short => {
665 let Some(keys) = one_rtt_keys.as_ref() else { continue };
666 if datagram.len() < 1 + server_scid.len() + 4 + 16 { continue; }
667 let parsed = LongPacket { packet_type: PacketType::Short, dcid: server_scid, scid: ConnectionId::from_slice(&[]), pn_offset: 1 + server_scid.len(), packet_end: datagram.len() };
668 let Ok((pn, payload)) = decrypt_payload(datagram, &parsed, &keys.remote) else { continue };
669 (payload, pn, nxtquic_crypto::Level::OneRtt)
670 }
671 _ => continue,
672 };
673 let mut frames = &payload[..];
674 while !frames.is_empty() {
675 let before = frames.len();
676 let Ok(frame) = Frame::decode(&mut frames) else { break };
677 match frame {
678 Frame::Crypto(crypto) if matches!(level, nxtquic_crypto::Level::Initial | nxtquic_crypto::Level::Handshake) => {
679 let Ok(output) = handshake.process(&crypto.data) else { continue };
680 while let Some((next_level, keys)) = handshake.next_keys() {
681 if next_level == nxtquic_crypto::Level::OneRtt { one_rtt_keys = Some(keys); } else { handshake_keys = Some((next_level, keys)); }
682 }
683 if !output.data.is_empty() {
684 let (packet_type, packet_number, send_keys) = match output.level {
685 nxtquic_crypto::Level::Initial => (PacketType::Initial, 1, Some(&initial_keys.server)),
686 nxtquic_crypto::Level::Handshake => (PacketType::Handshake, largest_handshake + 1, handshake_keys.as_ref().map(|(_, k)| &k.local)),
687 _ => (PacketType::Handshake, largest_handshake + 1, None),
688 };
689 let _ = send_protected_long(&socket, remote_addr, packet_type, initial.scid, server_scid, packet_number, &output.data, send_keys).await;
690 }
691 if handshake.is_complete() {
692 if let Some(keys) = one_rtt_keys.as_ref() {
693 let mut handshake_done = Vec::new();
694 Frame::HandshakeDone.encode(&mut handshake_done);
695 let _ = send_protected_short(&socket, remote_addr, initial.scid, 0, &handshake_done, &keys.local).await;
696 }
697 }
698 }
699 Frame::Stream(stream) if level == nxtquic_crypto::Level::OneRtt => {
700 let mut encoded = Vec::new();
701 Frame::Stream(stream).encode(&mut encoded);
702 let _ = connection.ingest_frames(&encoded).await;
703 }
704 _ => {}
705 }
706 if frames.len() >= before { break; }
707 }
708 let mut ack = Vec::new();
709 Frame::Ack(AckFrame {
710 largest_acknowledged: VarInt::from_u64(packet_number).unwrap(),
711 ack_delay: VarInt::ZERO,
712 first_ack_range: VarInt::ZERO,
713 ack_ranges: Vec::new(),
714 ecn_counts: None,
715 }).encode(&mut ack);
716 match level {
717 nxtquic_crypto::Level::Initial => {
718 let _ = send_protected_long(&socket, remote_addr, PacketType::Initial, initial.scid, server_scid, 2, &ack, Some(&initial_keys.server)).await;
719 }
720 nxtquic_crypto::Level::Handshake => {
721 if let Some((_, keys)) = handshake_keys.as_ref() {
722 let _ = send_protected_long(&socket, remote_addr, PacketType::Handshake, initial.scid, server_scid, largest_handshake + 1, &ack, Some(&keys.local)).await;
723 }
724 }
725 nxtquic_crypto::Level::OneRtt => {
726 if let Some(keys) = one_rtt_keys.as_ref() {
727 let _ = send_protected_short(&socket, remote_addr, initial.scid, application_packet_number, &ack, &keys.local).await;
728 application_packet_number += 1;
729 }
730 }
731 _ => {}
732 }
733 let _ = packet_number;
734 }
735 connection.close().await;
736}
737
738async fn send_protected_long(
739 socket: &tokio::net::UdpSocket,
740 remote_addr: SocketAddr,
741 packet_type: PacketType,
742 dcid: ConnectionId,
743 scid: ConnectionId,
744 packet_number: u64,
745 payload: &[u8],
746 keys: Option<&DirectionalKeys>,
747) -> io::Result<()> {
748 let Some(keys) = keys else { return Ok(()); };
749 let type_bits = match packet_type { PacketType::Handshake => 0x20, PacketType::Initial => 0x00, _ => return Ok(()) };
750 let mut header = vec![0xC0 | type_bits, 0, 0, 0, 1, dcid.len() as u8];
751 header.extend_from_slice(dcid.as_bytes());
752 header.push(scid.len() as u8);
753 header.extend_from_slice(scid.as_bytes());
754 if packet_type == PacketType::Initial { VarInt::ZERO.encode(&mut header); }
755 let mut body = payload.to_vec();
756 while body.len() < 20 {
757 body.push(0);
758 }
759 let packet_len = 1 + body.len() + keys.packet_key.tag_len();
760 VarInt::from_u64(packet_len as u64).map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "packet too large"))?.encode(&mut header);
761 let pn_offset = header.len();
762 header.push(packet_number as u8);
763 keys.packet_key.encrypt(packet_number, &header, &mut body).map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
764 header.extend_from_slice(&body);
765 let sample: [u8; 16] = header[pn_offset + 4..pn_offset + 20].try_into().map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "packet too short for sample"))?;
766 let mask = keys.header_key.protection_mask(&sample).map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
767 apply_header_protection(&mut header, &mask, pn_offset, 1);
768 socket.send_to(&header, remote_addr).await.map(|_| ())
769}
770
771async fn send_protected_short(
772 socket: &tokio::net::UdpSocket,
773 remote_addr: SocketAddr,
774 dcid: ConnectionId,
775 packet_number: u64,
776 payload: &[u8],
777 keys: &DirectionalKeys,
778) -> io::Result<()> {
779 let mut header = vec![0x40];
780 header.extend_from_slice(dcid.as_bytes());
781 let pn_offset = header.len();
782 header.push(packet_number as u8);
783 let mut body = payload.to_vec();
784 while body.len() < 20 {
785 body.push(0);
786 }
787 keys.packet_key.encrypt(packet_number, &header, &mut body).map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
788 header.extend_from_slice(&body);
789 let sample: [u8; 16] = header[pn_offset + 4..pn_offset + 20].try_into().map_err(|_| io::Error::new(io::ErrorKind::InvalidData, "packet too short for sample"))?;
790 let mask = keys.header_key.protection_mask(&sample).map_err(|error| io::Error::new(io::ErrorKind::Other, error))?;
791 apply_header_protection(&mut header, &mask, pn_offset, 1);
792 socket.send_to(&header, remote_addr).await.map(|_| ())
793}
794
795#[cfg(test)]
796mod packet_tests {
797 use super::*;
798 use bytes::Bytes;
799 use nxtquic_crypto::{CryptoProvider, RustlsCryptoProvider};
800 use nxtquic_proto::{
801 ConnectionId, Version,
802 frame::{CryptoFrame, Frame},
803 packet::protection::apply_header_protection,
804 varint::VarInt,
805 };
806
807 fn protected_client_initial() -> Vec<u8> {
808 let dcid = ConnectionId::from_slice(&[0x10; 8]);
809 let scid = [0x20; 8];
810 let provider = RustlsCryptoProvider::new(None, None);
811 let keys = provider.initial_keys(Version::V1, &dcid);
812
813 let mut plaintext = Vec::new();
814 Frame::Ping.encode(&mut plaintext);
815 Frame::Crypto(CryptoFrame {
816 offset: VarInt::ZERO,
817 data: Bytes::from_static(b"client hello bytes"),
818 })
819 .encode(&mut plaintext);
820
821 let ciphertext_len = plaintext.len() + keys.client.packet_key.tag_len();
823 let mut header = vec![0xc0, 0, 0, 0, 1, dcid.len() as u8];
824 header.extend_from_slice(dcid.as_bytes());
825 header.push(scid.len() as u8);
826 header.extend_from_slice(&scid);
827 header.push(0); VarInt::from_u64((1 + ciphertext_len) as u64)
829 .unwrap()
830 .encode(&mut header);
831 let packet_number_offset = header.len();
832 header.push(0); keys.client
835 .packet_key
836 .encrypt(0, &header, &mut plaintext)
837 .unwrap();
838 header.extend_from_slice(&plaintext);
839 let sample: [u8; 16] = header[packet_number_offset + 4..packet_number_offset + 20]
840 .try_into()
841 .unwrap();
842 let mask = keys.client.header_key.protection_mask(&sample).unwrap();
843 apply_header_protection(&mut header, &mask, packet_number_offset, 1);
844 header
845 }
846
847 #[test]
848 fn decrypts_client_initial_and_decodes_crypto_frames() {
849 let packet = protected_client_initial();
850 let initial = decrypt_client_initial(&packet).unwrap();
851 assert_eq!(initial.packet_number, 0);
852 assert!(matches!(initial.frames.first(), Some(Frame::Ping)));
853 assert!(
854 matches!(initial.frames.get(1), Some(Frame::Crypto(frame)) if frame.data == Bytes::from_static(b"client hello bytes"))
855 );
856 }
857
858 #[test]
859 fn rejects_tampered_client_initial() {
860 let mut packet = protected_client_initial();
861 let last = packet.len() - 1;
862 packet[last] ^= 0x01;
863 assert!(decrypt_client_initial(&packet).is_err());
864 }
865}