Skip to main content

wings_api/
tunnel.rs

1use super::client::{ApiHttpError, WingsClient};
2use futures_util::{SinkExt, StreamExt, ready};
3use std::{
4    io,
5    pin::Pin,
6    task::{Context, Poll},
7    time::Duration,
8};
9use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
10use tokio_tungstenite::{
11    MaybeTlsStream, WebSocketStream,
12    tungstenite::{Error as WsError, Message},
13};
14
15type WsStream = WebSocketStream<MaybeTlsStream<tokio::net::TcpStream>>;
16
17pub const MAX_DATAGRAM_SIZE: usize = 65536;
18
19const RECV_TIMEOUT: Duration = Duration::from_secs(5);
20
21const REFUSED_SIGNAL: &str = "refused";
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
24pub enum QueryProtocol {
25    Tcp,
26    Udp,
27}
28
29impl QueryProtocol {
30    fn as_str(self) -> &'static str {
31        match self {
32            QueryProtocol::Tcp => "tcp",
33            QueryProtocol::Udp => "udp",
34        }
35    }
36}
37
38impl WingsClient {
39    async fn open_tunnel(
40        &self,
41        server: uuid::Uuid,
42        protocol: QueryProtocol,
43        port: u16,
44    ) -> Result<WsStream, ApiHttpError> {
45        self.open_websocket(
46            format!(
47                "/api/servers/{server}/ws/query?protocol={}&port={port}",
48                protocol.as_str()
49            ),
50            reqwest::header::HeaderMap::new(),
51        )
52        .await
53    }
54
55    pub async fn open_tunnel_tcp(
56        &self,
57        server: uuid::Uuid,
58        port: u16,
59    ) -> Result<QueryTcpTunnel, ApiHttpError> {
60        Ok(QueryTcpTunnel {
61            stream: self.open_tunnel(server, QueryProtocol::Tcp, port).await?,
62            read: Vec::new(),
63            read_pos: 0,
64        })
65    }
66
67    pub async fn open_tunnel_udp(
68        &self,
69        server: uuid::Uuid,
70        port: u16,
71    ) -> Result<QueryUdpTunnel, ApiHttpError> {
72        Ok(QueryUdpTunnel {
73            stream: self.open_tunnel(server, QueryProtocol::Udp, port).await?,
74        })
75    }
76}
77
78pub struct QueryTcpTunnel {
79    stream: WsStream,
80    read: Vec<u8>,
81    read_pos: usize,
82}
83
84impl AsyncRead for QueryTcpTunnel {
85    fn poll_read(
86        mut self: Pin<&mut Self>,
87        cx: &mut Context<'_>,
88        buf: &mut ReadBuf<'_>,
89    ) -> Poll<io::Result<()>> {
90        loop {
91            if self.read_pos < self.read.len() {
92                let n = (self.read.len() - self.read_pos).min(buf.remaining());
93                let start = self.read_pos;
94                buf.put_slice(&self.read[start..start + n]);
95                self.read_pos += n;
96                return Poll::Ready(Ok(()));
97            }
98
99            match ready!(self.stream.poll_next_unpin(cx)) {
100                Some(Ok(Message::Binary(data))) => {
101                    self.read = data.to_vec();
102                    self.read_pos = 0;
103                }
104                Some(Ok(Message::Close(_))) | None => return Poll::Ready(Ok(())),
105                Some(Ok(_)) => {}
106                Some(Err(err)) => return Poll::Ready(Err(ws_to_io(err))),
107            }
108        }
109    }
110}
111
112impl AsyncWrite for QueryTcpTunnel {
113    fn poll_write(
114        mut self: Pin<&mut Self>,
115        cx: &mut Context<'_>,
116        buf: &[u8],
117    ) -> Poll<io::Result<usize>> {
118        ready!(self.stream.poll_ready_unpin(cx)).map_err(ws_to_io)?;
119        self.stream
120            .start_send_unpin(Message::binary(buf.to_vec()))
121            .map_err(ws_to_io)?;
122        Poll::Ready(Ok(buf.len()))
123    }
124
125    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
126        self.stream.poll_flush_unpin(cx).map_err(ws_to_io)
127    }
128
129    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
130        self.stream.poll_close_unpin(cx).map_err(ws_to_io)
131    }
132}
133
134pub struct QueryUdpTunnel {
135    stream: WsStream,
136}
137
138impl QueryUdpTunnel {
139    pub async fn send(&mut self, data: &[u8]) -> io::Result<()> {
140        self.stream
141            .send(Message::binary(data.to_vec()))
142            .await
143            .map_err(ws_to_io)
144    }
145
146    pub async fn recv(&mut self, buf: &mut [u8]) -> io::Result<usize> {
147        loop {
148            let message = match tokio::time::timeout(RECV_TIMEOUT, self.stream.next()).await {
149                Ok(Some(Ok(message))) => message,
150                Ok(Some(Err(err))) => return Err(ws_to_io(err)),
151                Ok(None) => return Err(io::ErrorKind::UnexpectedEof.into()),
152                Err(_) => return Err(io::ErrorKind::TimedOut.into()),
153            };
154
155            match message {
156                Message::Binary(data) => {
157                    let n = data.len().min(buf.len());
158                    buf[..n].copy_from_slice(&data[..n]);
159                    return Ok(n);
160                }
161                Message::Close(_) => return Err(io::ErrorKind::UnexpectedEof.into()),
162                Message::Text(text) if text.as_str() == REFUSED_SIGNAL => {
163                    return Err(io::ErrorKind::ConnectionRefused.into());
164                }
165                _ => {}
166            }
167        }
168    }
169}
170
171fn ws_to_io(err: WsError) -> io::Error {
172    match err {
173        WsError::Io(err) => err,
174        other => io::Error::other(other),
175    }
176}