mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-08-06 04:29:52 +00:00
fix(tunnel): preserve websocket cleanup after send errors (#2468)
Co-authored-by: FrankHan <frankhan@FrankHans-Mac-mini.local>
This commit is contained in:
@@ -10,10 +10,12 @@ use easytier_core::{
|
|||||||
tunnel::{IpVersion, Tunnel, TunnelError, wrapper::TunnelWrapper},
|
tunnel::{IpVersion, Tunnel, TunnelError, wrapper::TunnelWrapper},
|
||||||
};
|
};
|
||||||
use forwarded_header_value::ForwardedHeaderValue;
|
use forwarded_header_value::ForwardedHeaderValue;
|
||||||
use futures::{SinkExt, StreamExt};
|
use futures::{Sink, StreamExt};
|
||||||
use std::{
|
use std::{
|
||||||
net::{IpAddr, SocketAddr},
|
net::{IpAddr, SocketAddr},
|
||||||
|
pin::Pin,
|
||||||
sync::{Arc, LazyLock},
|
sync::{Arc, LazyLock},
|
||||||
|
task::{Context, Poll},
|
||||||
time::Duration,
|
time::Duration,
|
||||||
};
|
};
|
||||||
use tokio::{net::TcpListener, time::timeout};
|
use tokio::{net::TcpListener, time::timeout};
|
||||||
@@ -59,8 +61,46 @@ fn is_wss(url: &url::Url) -> Result<bool, TunnelError> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn sink_from_zc_packet<E>(packet: ZCPacket) -> Result<Message, E> {
|
struct WebSocketPacketSink<S> {
|
||||||
Ok(Message::binary(packet.tunnel_payload_bytes().freeze()))
|
inner: S,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<S> WebSocketPacketSink<S> {
|
||||||
|
fn new(inner: S) -> Self {
|
||||||
|
Self { inner }
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl<S, E> Sink<ZCPacket> for WebSocketPacketSink<S>
|
||||||
|
where
|
||||||
|
S: Sink<Message, Error = E> + Unpin,
|
||||||
|
E: std::fmt::Display,
|
||||||
|
{
|
||||||
|
type Error = TunnelError;
|
||||||
|
|
||||||
|
fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||||||
|
Pin::new(&mut self.inner)
|
||||||
|
.poll_ready(cx)
|
||||||
|
.map(|result| result.map_err(websocket_error))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn start_send(mut self: Pin<&mut Self>, packet: ZCPacket) -> Result<(), Self::Error> {
|
||||||
|
Pin::new(&mut self.inner)
|
||||||
|
.start_send(Message::binary(packet.tunnel_payload_bytes().freeze()))
|
||||||
|
.map_err(websocket_error)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||||||
|
Pin::new(&mut self.inner)
|
||||||
|
.poll_flush(cx)
|
||||||
|
.map(|result| result.map_err(websocket_error))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
|
||||||
|
Pin::new(&mut self.inner)
|
||||||
|
.poll_close(cx)
|
||||||
|
.map(|result| result.map_err(websocket_error))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn map_from_ws_message(
|
async fn map_from_ws_message(
|
||||||
@@ -231,9 +271,7 @@ where
|
|||||||
};
|
};
|
||||||
Ok(Box::new(TunnelWrapper::new(
|
Ok(Box::new(TunnelWrapper::new(
|
||||||
read.filter_map(map_from_ws_message),
|
read.filter_map(map_from_ws_message),
|
||||||
write
|
WebSocketPacketSink::new(write),
|
||||||
.sink_map_err(websocket_error)
|
|
||||||
.with(sink_from_zc_packet::<TunnelError>),
|
|
||||||
Some(info),
|
Some(info),
|
||||||
)))
|
)))
|
||||||
}
|
}
|
||||||
@@ -365,9 +403,7 @@ where
|
|||||||
let (write, read) = client.split();
|
let (write, read) = client.split();
|
||||||
Ok(Box::new(TunnelWrapper::new(
|
Ok(Box::new(TunnelWrapper::new(
|
||||||
read.filter_map(map_from_ws_message),
|
read.filter_map(map_from_ws_message),
|
||||||
write
|
WebSocketPacketSink::new(write),
|
||||||
.sink_map_err(websocket_error)
|
|
||||||
.with(sink_from_zc_packet::<TunnelError>),
|
|
||||||
Some(info),
|
Some(info),
|
||||||
)))
|
)))
|
||||||
}
|
}
|
||||||
@@ -376,11 +412,66 @@ where
|
|||||||
pub mod tests {
|
pub mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
use easytier_core::socket::SocketListener;
|
use easytier_core::socket::SocketListener;
|
||||||
|
use futures::SinkExt;
|
||||||
|
use std::io;
|
||||||
use tokio::{
|
use tokio::{
|
||||||
io::{AsyncReadExt, AsyncWriteExt},
|
io::{AsyncReadExt, AsyncWriteExt},
|
||||||
net::TcpSocket,
|
net::TcpSocket,
|
||||||
};
|
};
|
||||||
|
|
||||||
|
struct FailingWebSocketSink {
|
||||||
|
close_called: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Sink<Message> for FailingWebSocketSink {
|
||||||
|
type Error = io::Error;
|
||||||
|
|
||||||
|
fn poll_ready(
|
||||||
|
self: Pin<&mut Self>,
|
||||||
|
_cx: &mut Context<'_>,
|
||||||
|
) -> Poll<Result<(), Self::Error>> {
|
||||||
|
Poll::Ready(Ok(()))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn start_send(self: Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn poll_flush(
|
||||||
|
self: Pin<&mut Self>,
|
||||||
|
_cx: &mut Context<'_>,
|
||||||
|
) -> Poll<Result<(), Self::Error>> {
|
||||||
|
Poll::Ready(Err(io::Error::new(
|
||||||
|
io::ErrorKind::BrokenPipe,
|
||||||
|
"send failed",
|
||||||
|
)))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn poll_close(
|
||||||
|
mut self: Pin<&mut Self>,
|
||||||
|
_cx: &mut Context<'_>,
|
||||||
|
) -> Poll<Result<(), Self::Error>> {
|
||||||
|
self.close_called = true;
|
||||||
|
Poll::Ready(Err(io::Error::new(
|
||||||
|
io::ErrorKind::BrokenPipe,
|
||||||
|
"close failed",
|
||||||
|
)))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn packet_sink_maps_send_and_close_errors_independently() {
|
||||||
|
let mut sink = WebSocketPacketSink::new(FailingWebSocketSink {
|
||||||
|
close_called: false,
|
||||||
|
});
|
||||||
|
|
||||||
|
sink.send(ZCPacket::new_with_payload(b"packet"))
|
||||||
|
.await
|
||||||
|
.expect_err("send should fail");
|
||||||
|
sink.close().await.expect_err("close should fail");
|
||||||
|
assert!(sink.inner.close_called);
|
||||||
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn ws_forwarded() {
|
async fn ws_forwarded() {
|
||||||
let mut listener = WsTunnelListener::new("ws://127.0.0.1:25559".parse().unwrap());
|
let mut listener = WsTunnelListener::new("ws://127.0.0.1:25559".parse().unwrap());
|
||||||
|
|||||||
Reference in New Issue
Block a user