ant_quic/masque/
relay_socket.rs1use 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
31pub struct MasqueRelaySocket {
34 relay_public_addr: SocketAddr,
36 recv_queue: std::sync::Mutex<VecDeque<(Vec<u8>, SocketAddr)>>,
38 recv_waker: std::sync::Mutex<Option<Waker>>,
40 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 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 let socket_ref = Arc::clone(&socket);
78 tokio::spawn(async move {
79 loop {
80 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 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 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; 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 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 tokio::spawn(async move {
131 loop {
132 let Some(encoded) = send_rx.recv().await else {
133 if let Err(e) = send_stream.finish() {
138 tracing::debug!(error = %e, "MasqueRelaySocket: stream finish error");
139 }
140 return;
141 };
142
143 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 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}