Skip to main content

ant_quic/masque/
relay_socket.rs

1// Copyright 2024 Saorsa Labs Ltd.
2//
3// This Saorsa Network Software is licensed under the General Public License (GPL), version 3.
4// Please see the file LICENSE-GPL, or visit <http://www.gnu.org/licenses/> for the full text.
5//
6// Full details available at https://saorsalabs.com/licenses
7
8//! MASQUE Relay Socket
9//!
10//! A virtual UDP socket that routes QUIC packets through a MASQUE relay
11//! via a persistent QUIC stream (length-prefixed framing).
12//!
13//! Implements [`AsyncUdpSocket`] so it can be plugged into a Quinn endpoint
14//! as a transparent replacement for a real UDP socket.
15
16use bytes::Bytes;
17use std::collections::VecDeque;
18use std::fmt;
19use std::io::{self, IoSliceMut};
20use std::net::SocketAddr;
21use std::pin::Pin;
22use std::sync::Arc;
23use std::task::{Context, Poll, Waker};
24
25use quinn_udp::{RecvMeta, Transmit};
26
27use crate::VarInt;
28use crate::high_level::{AsyncUdpSocket, UdpSender};
29use crate::masque::UncompressedDatagram;
30
31/// A virtual UDP socket that tunnels packets through a MASQUE relay
32/// via a persistent QUIC stream with length-prefixed framing.
33pub struct MasqueRelaySocket {
34    /// The relay's public address (returned as our local address)
35    relay_public_addr: SocketAddr,
36    /// Queue of received packets (payload, source_addr)
37    recv_queue: std::sync::Mutex<VecDeque<(Vec<u8>, SocketAddr)>>,
38    /// Waker to notify when new packets arrive
39    recv_waker: std::sync::Mutex<Option<Waker>>,
40    /// Channel for outbound packets (written to the relay stream by background task)
41    send_tx: tokio::sync::mpsc::UnboundedSender<Bytes>,
42}
43
44impl fmt::Debug for MasqueRelaySocket {
45    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
46        f.debug_struct("MasqueRelaySocket")
47            .field("relay_public_addr", &self.relay_public_addr)
48            .field(
49                "recv_queue_len",
50                &self.recv_queue.lock().map(|q| q.len()).unwrap_or(0),
51            )
52            .finish()
53    }
54}
55
56impl MasqueRelaySocket {
57    /// Create a new stream-based relay socket.
58    ///
59    /// Spawns two background tasks:
60    /// - Read from `recv_stream`, decode frames, queue for `poll_recv`
61    /// - Read from `send_tx` channel, write length-prefixed frames to `send_stream`
62    pub fn new(
63        mut send_stream: crate::high_level::SendStream,
64        mut recv_stream: crate::high_level::RecvStream,
65        relay_public_addr: SocketAddr,
66    ) -> Arc<Self> {
67        let (send_tx, mut send_rx) = tokio::sync::mpsc::unbounded_channel::<Bytes>();
68
69        let socket = Arc::new(Self {
70            relay_public_addr,
71            recv_queue: std::sync::Mutex::new(VecDeque::new()),
72            recv_waker: std::sync::Mutex::new(None),
73            send_tx,
74        });
75
76        // Background task: read length-prefixed frames from relay stream → queue
77        let socket_ref = Arc::clone(&socket);
78        tokio::spawn(async move {
79            loop {
80                // Read 4-byte length prefix
81                let mut len_buf = [0u8; 4];
82                if let Err(e) = recv_stream.read_exact(&mut len_buf).await {
83                    tracing::debug!(error = %e, "MasqueRelaySocket: stream read error (length)");
84                    break;
85                }
86                let frame_len = u32::from_be_bytes(len_buf) as usize;
87                if frame_len > 65536 {
88                    tracing::warn!(frame_len, "MasqueRelaySocket: oversized frame");
89                    break;
90                }
91
92                // Read frame data
93                let mut frame_buf = vec![0u8; frame_len];
94                if let Err(e) = recv_stream.read_exact(&mut frame_buf).await {
95                    tracing::debug!(error = %e, "MasqueRelaySocket: stream read error (data)");
96                    break;
97                }
98
99                // Decode as UncompressedDatagram
100                let mut cursor = Bytes::from(frame_buf);
101                match UncompressedDatagram::decode(&mut cursor) {
102                    Ok(datagram) => {
103                        let payload = datagram.payload.to_vec();
104                        let source = datagram.target; // "target" in datagram = source from relay's perspective
105
106                        if let Ok(mut queue) = socket_ref.recv_queue.lock() {
107                            queue.push_back((payload, source));
108                        }
109                        if let Ok(mut waker) = socket_ref.recv_waker.lock() {
110                            if let Some(w) = waker.take() {
111                                w.wake();
112                            }
113                        }
114                    }
115                    Err(_) => {
116                        tracing::trace!("MasqueRelaySocket: failed to decode frame");
117                    }
118                }
119            }
120
121            // Wake pending recv on stream close
122            if let Ok(mut waker) = socket_ref.recv_waker.lock() {
123                if let Some(w) = waker.take() {
124                    w.wake();
125                }
126            }
127        });
128
129        // Background task: write queued outbound packets to relay stream
130        tokio::spawn(async move {
131            loop {
132                let Some(encoded) = send_rx.recv().await else {
133                    // Every sender is gone, i.e. the socket was dropped. Nothing is
134                    // half-written, so finish gracefully: QUIC keeps retransmitting
135                    // frames that are queued but not yet acknowledged, so packets
136                    // already handed to the relay still arrive.
137                    if let Err(e) = send_stream.finish() {
138                        tracing::debug!(error = %e, "MasqueRelaySocket: stream finish error");
139                    }
140                    return;
141                };
142
143                // A failed write may have left a partial frame on the wire. Return
144                // without finishing so the drop resets the stream, and the relay's
145                // length-prefixed reader sees a stream error rather than misparsing
146                // a truncated frame.
147                let frame_len = encoded.len() as u32;
148                if let Err(e) = send_stream.write_all(&frame_len.to_be_bytes()).await {
149                    tracing::debug!(error = %e, "MasqueRelaySocket: stream write error (length)");
150                    return;
151                }
152                if let Err(e) = send_stream.write_all(&encoded).await {
153                    tracing::debug!(error = %e, "MasqueRelaySocket: stream write error (data)");
154                    return;
155                }
156            }
157        });
158
159        socket
160    }
161}
162
163impl AsyncUdpSocket for MasqueRelaySocket {
164    fn create_sender(&self) -> Pin<Box<dyn UdpSender>> {
165        Box::pin(MasqueRelaySender {
166            send_tx: self.send_tx.clone(),
167        })
168    }
169
170    fn poll_recv(
171        &self,
172        cx: &mut Context,
173        bufs: &mut [IoSliceMut<'_>],
174        meta: &mut [RecvMeta],
175    ) -> Poll<io::Result<usize>> {
176        if bufs.is_empty() || meta.is_empty() {
177            return Poll::Ready(Ok(0));
178        }
179
180        if let Ok(mut queue) = self.recv_queue.lock() {
181            if let Some((payload, source)) = queue.pop_front() {
182                let len = payload.len().min(bufs[0].len());
183                bufs[0][..len].copy_from_slice(&payload[..len]);
184
185                let mut recv_meta = RecvMeta::default();
186                recv_meta.len = len;
187                recv_meta.stride = len;
188                recv_meta.addr = source;
189                recv_meta.ecn = None;
190                recv_meta.dst_ip = None;
191                meta[0] = recv_meta;
192
193                return Poll::Ready(Ok(1));
194            }
195        }
196
197        // Register waker for when data arrives
198        if let Ok(mut waker) = self.recv_waker.lock() {
199            *waker = Some(cx.waker().clone());
200        }
201
202        Poll::Pending
203    }
204
205    fn local_addr(&self) -> io::Result<SocketAddr> {
206        Ok(self.relay_public_addr)
207    }
208
209    fn may_fragment(&self) -> bool {
210        false
211    }
212}
213
214#[derive(Debug)]
215struct MasqueRelaySender {
216    send_tx: tokio::sync::mpsc::UnboundedSender<Bytes>,
217}
218
219impl UdpSender for MasqueRelaySender {
220    fn poll_send(
221        self: Pin<&mut Self>,
222        transmit: &Transmit,
223        _cx: &mut Context<'_>,
224    ) -> Poll<io::Result<()>> {
225        let datagram = UncompressedDatagram::new(
226            VarInt::from_u32(0),
227            transmit.destination,
228            Bytes::copy_from_slice(transmit.contents),
229        );
230        let encoded = datagram.encode();
231
232        Poll::Ready(
233            self.send_tx.send(encoded).map_err(|_| {
234                io::Error::new(io::ErrorKind::ConnectionAborted, "relay stream closed")
235            }),
236        )
237    }
238}