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}