use std::any::Any;
use std::collections::HashMap;
use std::fmt;
use std::io::{self, Read, Write};
use std::net::{TcpStream, ToSocketAddrs};
use std::panic::{catch_unwind, AssertUnwindSafe};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::mpsc::{self, Receiver, Sender, TryRecvError};
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
use std::thread;
use std::time::{Duration, Instant};
use tungstenite::client::IntoClientRequest;
use tungstenite::handshake::HandshakeError;
use tungstenite::protocol::frame::coding::CloseCode;
use tungstenite::protocol::{CloseFrame, WebSocketConfig};
use tungstenite::{Message, WebSocket};
use super::transport::{WsHandshake, WsLinkEvent, WsLinkId, WsTransport};
use super::WsFrame;
use crate::config::MAX_TIMEOUT;
use crate::response::{BackendError, RawResponse};
type Event = (WsLinkId, WsLinkEvent);
enum Command {
Send(WsFrame),
Close(u16),
}
const EVENT_BUDGET_MESSAGES: usize = 32;
fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
mutex.lock().unwrap_or_else(PoisonError::into_inner)
}
fn deadline_after(duration: Duration) -> Instant {
let now = Instant::now();
now.checked_add(duration.min(MAX_TIMEOUT)).unwrap_or(now)
}
struct LinkHandle {
commands: Sender<Command>,
queued_bytes: Arc<AtomicUsize>,
}
pub struct TungsteniteTransport {
links: HashMap<WsLinkId, LinkHandle>,
events_tx: Sender<Event>,
events_rx: Mutex<Receiver<Event>>,
immediate: Vec<Event>,
}
impl fmt::Debug for TungsteniteTransport {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("TungsteniteTransport").field("links", &self.links.len()).finish_non_exhaustive()
}
}
impl Default for TungsteniteTransport {
fn default() -> Self {
Self::new()
}
}
impl TungsteniteTransport {
pub fn new() -> Self {
let (events_tx, events_rx) = mpsc::channel();
Self { links: HashMap::new(), events_tx, events_rx: Mutex::new(events_rx), immediate: Vec::new() }
}
}
impl WsTransport for TungsteniteTransport {
fn open(&mut self, link: WsLinkId, handshake: WsHandshake) {
let (commands_tx, commands_rx) = mpsc::channel();
let events = self.events_tx.clone();
let queued_bytes = Arc::new(AtomicUsize::new(0));
let counter = Arc::clone(&queued_bytes);
let spawned = thread::Builder::new().name(format!("net-backend-ws-{link}")).spawn(move || {
let sink = EventSink { link, events, queued_bytes: counter };
let last = catch_unwind(AssertUnwindSafe(|| session(&handshake, &commands_rx, &sink, None)))
.unwrap_or_else(|panic| WsLinkEvent::Failed(BackendError::Network(format!("the WebSocket thread panicked: {}", panic_text(panic.as_ref())))));
let _ = sink.events.send((link, last));
});
match spawned {
Ok(_) => {
self.links.insert(link, LinkHandle { commands: commands_tx, queued_bytes });
}
Err(e) => self.immediate.push((link, WsLinkEvent::Failed(BackendError::Network(format!("could not start a WebSocket thread: {e}"))))),
}
}
fn send(&mut self, link: WsLinkId, frame: WsFrame) {
if let Some(handle) = self.links.get(&link) {
let _ = handle.commands.send(Command::Send(frame));
}
}
fn close(&mut self, link: WsLinkId, code: u16) {
if let Some(handle) = self.links.remove(&link) {
let _ = handle.commands.send(Command::Close(code));
}
}
fn poll(&mut self) -> Vec<Event> {
let mut out = std::mem::take(&mut self.immediate);
let rx = lock(&self.events_rx);
while let Ok(event) = rx.try_recv() {
out.push(event);
}
drop(rx);
for (link, event) in &out {
match event {
WsLinkEvent::Frame(frame) => {
if let Some(handle) = self.links.get(link) {
let _ = handle.queued_bytes.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| Some(n.saturating_sub(frame.len())));
}
}
WsLinkEvent::Closed { .. } | WsLinkEvent::Failed(_) => {
self.links.remove(link);
}
_ => {}
}
}
out
}
fn shutdown(&mut self) {
for (_, handle) in self.links.drain() {
let _ = handle.commands.send(Command::Close(1001));
}
}
}
impl Drop for TungsteniteTransport {
fn drop(&mut self) {
self.shutdown();
}
}
fn panic_text(panic: &(dyn Any + Send)) -> &str {
panic.downcast_ref::<&str>().copied().or_else(|| panic.downcast_ref::<String>().map(String::as_str)).unwrap_or("no message")
}
struct EventSink {
link: WsLinkId,
events: Sender<Event>,
queued_bytes: Arc<AtomicUsize>,
}
impl EventSink {
fn send(&self, event: WsLinkEvent) -> bool {
self.events.send((self.link, event)).is_ok()
}
}
struct TimedTcp {
tcp: TcpStream,
deadline: Option<Instant>,
max_wait: Duration,
applied: Option<Duration>,
last_byte: Instant,
}
impl TimedTcp {
fn new(tcp: TcpStream, deadline: Option<Instant>, max_wait: Duration) -> Self {
Self { tcp, deadline, max_wait, applied: None, last_byte: Instant::now() }
}
}
impl Read for TimedTcp {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
let mut wait = self.max_wait;
if let Some(deadline) = self.deadline {
let left = deadline.saturating_duration_since(Instant::now());
if left.is_zero() {
return Err(io::Error::new(io::ErrorKind::TimedOut, "read deadline reached"));
}
wait = wait.min(left);
}
let wait = wait.max(Duration::from_millis(1));
if self.applied != Some(wait) {
self.tcp.set_read_timeout(Some(wait))?;
self.applied = Some(wait);
}
let n = loop {
match self.tcp.read(buf) {
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
other => break other?,
}
};
if n > 0 {
self.last_byte = Instant::now();
}
Ok(n)
}
}
impl Write for TimedTcp {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
self.tcp.write(buf)
}
fn flush(&mut self) -> io::Result<()> {
self.tcp.flush()
}
}
enum Stream {
Plain(TimedTcp),
Tls(Box<rustls::StreamOwned<rustls::ClientConnection, TimedTcp>>),
}
impl Stream {
fn timed(&self) -> &TimedTcp {
match self {
Stream::Plain(tcp) => tcp,
Stream::Tls(tls) => &tls.sock,
}
}
fn timed_mut(&mut self) -> &mut TimedTcp {
match self {
Stream::Plain(tcp) => tcp,
Stream::Tls(tls) => &mut tls.sock,
}
}
}
impl Read for Stream {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
match self {
Stream::Plain(tcp) => tcp.read(buf),
Stream::Tls(tls) => tls.read(buf),
}
}
}
impl Write for Stream {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
match self {
Stream::Plain(tcp) => tcp.write(buf),
Stream::Tls(tls) => tls.write(buf),
}
}
fn flush(&mut self) -> io::Result<()> {
match self {
Stream::Plain(tcp) => tcp.flush(),
Stream::Tls(tls) => tls.flush(),
}
}
}
fn is_no_data(error: &io::Error) -> bool {
matches!(error.kind(), io::ErrorKind::WouldBlock | io::ErrorKind::TimedOut)
}
fn map_error(error: tungstenite::Error) -> BackendError {
match error {
tungstenite::Error::Http(response) => {
let (parts, body) = (*response).into_parts();
BackendError::Status(Box::new(RawResponse { status: parts.status, headers: parts.headers, body: body.unwrap_or_default() }))
}
tungstenite::Error::Io(ref io) if is_no_data(io) => BackendError::Timeout(format!("socket: {io}")),
tungstenite::Error::Io(ref io) if io.get_ref().is_some_and(|inner| inner.is::<rustls::Error>()) => BackendError::Tls(error.to_string()),
tungstenite::Error::Capacity(_) => BackendError::disconnected("a message was larger than the limit (closed with 1009)", None),
tungstenite::Error::Url(_) | tungstenite::Error::HttpFormat(_) => BackendError::InvalidRequest(error.to_string()),
_ => BackendError::Network(error.to_string()),
}
}
fn host_for_connect(uri: &http::Uri) -> Option<String> {
let host = uri.host()?;
Some(host.strip_prefix('[').and_then(|h| h.strip_suffix(']')).unwrap_or(host).to_string())
}
fn close_requested(commands: &Receiver<Command>) -> bool {
matches!(commands.try_recv(), Ok(Command::Close(_)) | Err(TryRecvError::Disconnected))
}
fn connect(handshake: &WsHandshake, commands: &Receiver<Command>, tls: Option<Arc<rustls::ClientConfig>>) -> Result<Option<WebSocket<Stream>>, BackendError> {
let deadline = deadline_after(handshake.connect_timeout);
let host = host_for_connect(&handshake.uri).ok_or_else(|| BackendError::InvalidRequest("the URL has no host".into()))?;
let secure = handshake.is_secure();
let port = handshake.uri.port_u16().unwrap_or(if secure { 443 } else { 80 });
let addrs: Vec<_> = (host.as_str(), port).to_socket_addrs().map_err(|e| BackendError::Network(format!("could not resolve the host: {e}")))?.collect();
let mut last = BackendError::Network("the host resolved to no address".into());
let mut tcp = None;
for addr in addrs {
let left = deadline.saturating_duration_since(Instant::now());
if left.is_zero() {
last = BackendError::Timeout("connect limit".into());
break;
}
match TcpStream::connect_timeout(&addr, left) {
Ok(stream) => {
tcp = Some(stream);
break;
}
Err(e) if e.kind() == io::ErrorKind::TimedOut => last = BackendError::Timeout("connect limit".into()),
Err(e) => last = BackendError::Network(format!("could not connect: {e}")),
}
}
let tcp = tcp.ok_or(last)?;
let io_error = |e: io::Error| BackendError::Network(format!("socket setup: {e}"));
tcp.set_nodelay(true).map_err(io_error)?;
tcp.set_write_timeout(Some(write_limit(handshake))).map_err(io_error)?;
let timed = TimedTcp::new(tcp, Some(deadline), handshake.connect_timeout.min(MAX_TIMEOUT));
let stream = if secure {
let config = match tls {
Some(config) => config,
None => crate::tls::client_config().map_err(BackendError::Tls)?,
};
let name = rustls::pki_types::ServerName::try_from(host.clone()).map_err(|e| BackendError::InvalidRequest(format!("bad TLS server name: {e}")))?;
let connection = rustls::ClientConnection::new(config, name).map_err(|e| BackendError::Tls(e.to_string()))?;
Stream::Tls(Box::new(rustls::StreamOwned::new(connection, timed)))
} else {
Stream::Plain(timed)
};
let mut request = handshake.uri.clone().into_client_request().map_err(map_error)?;
const WS_HEADERS: [&str; 5] = ["host", "connection", "upgrade", "sec-websocket-version", "sec-websocket-key"];
for (name, value) in &handshake.headers {
if !WS_HEADERS.contains(&name.as_str()) {
request.headers_mut().append(name.clone(), value.clone());
}
}
if close_requested(commands) {
return Ok(None);
}
let config = WebSocketConfig::default().max_message_size(Some(handshake.max_message_bytes)).max_frame_size(Some(handshake.max_message_bytes));
let past_deadline = |error: BackendError| if Instant::now() >= deadline { BackendError::Timeout("connect limit".into()) } else { error };
let mut attempt = tungstenite::client::client_with_config(request, stream, Some(config));
let mut ws = loop {
match attempt {
Ok((ws, _response)) => break ws,
Err(HandshakeError::Failure(e)) => return Err(past_deadline(map_error(e))),
Err(HandshakeError::Interrupted(mid)) => {
if Instant::now() >= deadline {
return Err(BackendError::Timeout("connect limit".into()));
}
attempt = mid.handshake();
}
}
};
let timed = ws.get_mut().timed_mut();
timed.deadline = None;
timed.max_wait = handshake.read_timeout;
Ok(Some(ws))
}
fn to_message(frame: WsFrame) -> Message {
match frame {
WsFrame::Text(text) => Message::text(text),
WsFrame::Binary(bytes) => Message::binary(bytes),
}
}
fn close_with(ws: &mut WebSocket<Stream>, code: CloseCode) {
let _ = ws.close(Some(CloseFrame { code, reason: "".into() }));
}
fn session(handshake: &WsHandshake, commands: &Receiver<Command>, sink: &EventSink, tls: Option<Arc<rustls::ClientConfig>>) -> WsLinkEvent {
let closed_by_game = || WsLinkEvent::Closed { code: None, reason: "closed by the game".into() };
let mut ws = match connect(handshake, commands, tls) {
Ok(Some(ws)) => ws,
Ok(None) => return closed_by_game(),
Err(error) => return WsLinkEvent::Failed(error),
};
if close_requested(commands) {
close_with(&mut ws, CloseCode::Normal);
let _ = ws.flush();
return closed_by_game();
}
if !sink.send(WsLinkEvent::Opened) {
return WsLinkEvent::Closed { code: None, reason: "the plugin is gone".into() };
}
let budget = EVENT_BUDGET_MESSAGES.saturating_mul(handshake.max_message_bytes);
let write_limit = write_limit(handshake);
let mut last_ping = Instant::now();
let mut alive = Instant::now();
let mut closing: Option<Instant> = None;
let mut failure: Option<BackendError> = None;
let mut close_code: Option<u16> = None;
let mut close_reason = String::new();
let finish = |failure: Option<BackendError>, code: Option<u16>, reason: String| match failure {
Some(error) => WsLinkEvent::Failed(error),
None => WsLinkEvent::Closed { code, reason },
};
let write_failed = |e: tungstenite::Error| match e {
tungstenite::Error::Io(ref io) if is_no_data(io) => BackendError::Timeout(format!("the server accepted no data for {write_limit:?}")),
other => map_error(other),
};
loop {
let before_writes = Instant::now();
loop {
match commands.try_recv() {
Ok(Command::Send(frame)) if closing.is_none() => {
if let Err(e) = ws.write(to_message(frame)) {
return WsLinkEvent::Failed(write_failed(e));
}
}
Ok(Command::Send(_)) => {}
Ok(Command::Close(code)) => {
if closing.is_none() {
close_with(&mut ws, CloseCode::from(code));
closing = Some(Instant::now());
}
}
Err(TryRecvError::Empty) => break,
Err(TryRecvError::Disconnected) => {
if closing.is_none() {
close_with(&mut ws, CloseCode::Away);
closing = Some(Instant::now());
}
break;
}
}
}
if closing.is_none() && last_ping.elapsed() >= handshake.ping_interval {
if let Err(e) = ws.write(Message::Ping(Default::default())) {
return WsLinkEvent::Failed(write_failed(e));
}
last_ping = Instant::now();
}
match ws.flush() {
Ok(()) => {}
Err(tungstenite::Error::ConnectionClosed | tungstenite::Error::AlreadyClosed) => return finish(failure, close_code, close_reason),
Err(_) if closing.is_some() => return finish(failure, close_code, close_reason),
Err(e) => return WsLinkEvent::Failed(write_failed(e)),
}
let blocked = before_writes.elapsed();
if blocked > handshake.read_timeout {
alive = alive.checked_add(blocked).map_or_else(Instant::now, |a| a.min(Instant::now()));
}
ws.get_mut().timed_mut().deadline = Some(deadline_after(handshake.read_timeout));
match ws.read() {
Ok(message) => {
alive = Instant::now();
let frame = match message {
Message::Text(text) => Some(WsFrame::Text(text.as_str().to_string())),
Message::Binary(bytes) => Some(WsFrame::Binary(bytes.to_vec())),
Message::Close(frame) => {
if let Some(frame) = frame {
close_code = Some(u16::from(frame.code));
close_reason = frame.reason.as_str().to_string();
}
closing.get_or_insert_with(Instant::now);
None
}
_ => None,
};
if let (Some(frame), None) = (frame, &failure) {
let len = frame.len();
if sink.queued_bytes.load(Ordering::SeqCst).saturating_add(len) > budget {
close_with(&mut ws, CloseCode::Policy);
closing.get_or_insert_with(Instant::now);
failure = Some(BackendError::disconnected("the game did not take the received frames fast enough (closed with 1008)", None));
} else {
sink.queued_bytes.fetch_add(len, Ordering::SeqCst);
if !sink.send(WsLinkEvent::Frame(frame)) && closing.is_none() {
close_with(&mut ws, CloseCode::Away);
closing = Some(Instant::now());
}
}
}
}
Err(tungstenite::Error::Io(ref e)) if is_no_data(e) => {}
Err(tungstenite::Error::ConnectionClosed | tungstenite::Error::AlreadyClosed) => return finish(failure, close_code, close_reason),
Err(_) if closing.is_some() => return finish(failure, close_code, close_reason),
Err(tungstenite::Error::Capacity(_)) => {
close_with(&mut ws, CloseCode::Size);
let _ = ws.flush();
drain_below(&mut ws);
return WsLinkEvent::Failed(BackendError::disconnected("a message was larger than the limit (closed with 1009)", None));
}
Err(e) => return WsLinkEvent::Failed(map_error(e)),
}
if let Some(started) = closing {
if started.elapsed() > Duration::from_secs(1) {
return finish(failure, close_code, close_reason);
}
} else {
let last = alive.max(ws.get_ref().timed().last_byte);
if last.elapsed() > handshake.dead_after {
return WsLinkEvent::Failed(BackendError::Timeout(format!(
"no byte from the server within {:?}; connection considered dead",
handshake.dead_after
)));
}
}
}
}
fn drain_below(ws: &mut WebSocket<Stream>) {
let until = deadline_after(Duration::from_secs(1));
let stream = ws.get_mut();
stream.timed_mut().deadline = Some(until);
let mut scratch = [0u8; 16 * 1024];
while Instant::now() < until {
match stream.read(&mut scratch) {
Ok(0) => break,
Ok(_) => {}
Err(e) if is_no_data(&e) => {}
Err(_) => break,
}
}
}
fn write_limit(handshake: &WsHandshake) -> Duration {
handshake.dead_after.max(Duration::from_secs(30)).min(MAX_TIMEOUT)
}
#[cfg(test)]
mod tests {
use std::net::TcpListener;
use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer};
use super::*;
struct Pki {
ca: CertificateDer<'static>,
leaf: CertificateDer<'static>,
key: PrivateKeyDer<'static>,
}
fn pki() -> Pki {
let ca_key = rcgen::KeyPair::generate().unwrap_or_else(|e| panic!("{e}"));
let mut ca_params = rcgen::CertificateParams::new(Vec::<String>::new()).unwrap_or_else(|e| panic!("{e}"));
ca_params.is_ca = rcgen::IsCa::Ca(rcgen::BasicConstraints::Unconstrained);
let ca = rcgen::CertifiedIssuer::self_signed(ca_params, ca_key).unwrap_or_else(|e| panic!("{e}"));
let leaf_key = rcgen::KeyPair::generate().unwrap_or_else(|e| panic!("{e}"));
let leaf = rcgen::CertificateParams::new(vec!["localhost".to_string()])
.unwrap_or_else(|e| panic!("{e}"))
.signed_by(&leaf_key, &ca)
.unwrap_or_else(|e| panic!("{e}"));
Pki { ca: ca.der().clone(), leaf: leaf.der().clone(), key: PrivateKeyDer::Pkcs8(PrivatePkcs8KeyDer::from(leaf_key.serialize_der())) }
}
fn client_config(pki: &Pki) -> Arc<rustls::ClientConfig> {
let mut roots = rustls::RootCertStore::empty();
roots.add(pki.ca.clone()).unwrap_or_else(|e| panic!("{e}"));
let config = rustls::ClientConfig::builder_with_provider(crate::tls::provider())
.with_safe_default_protocol_versions()
.unwrap_or_else(|e| panic!("{e}"))
.with_root_certificates(roots)
.with_no_client_auth();
Arc::new(config)
}
fn server_config(pki: &Pki) -> Arc<rustls::ServerConfig> {
let config = rustls::ServerConfig::builder_with_provider(crate::tls::provider())
.with_safe_default_protocol_versions()
.unwrap_or_else(|e| panic!("{e}"))
.with_no_client_auth()
.with_single_cert(vec![pki.leaf.clone()], pki.key.clone_key())
.unwrap_or_else(|e| panic!("{e}"));
Arc::new(config)
}
fn handshake(uri: String, connect_timeout: Duration, read_timeout: Duration, ping: Duration, dead: Duration) -> WsHandshake {
WsHandshake {
uri: uri.parse().unwrap_or_else(|e| panic!("{e}")),
headers: http::HeaderMap::new(),
connect_timeout,
read_timeout,
ping_interval: ping,
dead_after: dead,
max_message_bytes: 4 << 20,
}
}
fn start(handshake: WsHandshake, tls: Option<Arc<rustls::ClientConfig>>) -> (Sender<Command>, Receiver<Event>, thread::JoinHandle<WsLinkEvent>) {
let (commands_tx, commands_rx) = mpsc::channel();
let (events_tx, events_rx) = mpsc::channel();
let sink = EventSink { link: WsLinkId::next(), events: events_tx, queued_bytes: Arc::default() };
let thread = thread::spawn(move || session(&handshake, &commands_rx, &sink, tls));
(commands_tx, events_rx, thread)
}
fn echo_server(pki: &Pki) -> (u16, thread::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").unwrap_or_else(|e| panic!("{e}"));
let port = listener.local_addr().map(|a| a.port()).unwrap_or_else(|e| panic!("{e}"));
let config = server_config(pki);
let handle = thread::spawn(move || {
let (tcp, _) = listener.accept().unwrap_or_else(|e| panic!("{e}"));
tcp.set_read_timeout(Some(Duration::from_secs(10))).unwrap_or_else(|e| panic!("{e}"));
let connection = rustls::ServerConnection::new(config).unwrap_or_else(|e| panic!("{e}"));
let tls = rustls::StreamOwned::new(connection, tcp);
let limits = WebSocketConfig::default().max_message_size(Some(8 << 20)).max_frame_size(Some(8 << 20));
let mut ws = tungstenite::accept_with_config(tls, Some(limits)).unwrap_or_else(|e| panic!("server handshake: {e}"));
ws.get_ref().get_ref().set_read_timeout(Some(Duration::from_millis(5))).unwrap_or_else(|e| panic!("{e}"));
let deadline = Instant::now() + Duration::from_secs(60);
while Instant::now() < deadline {
match ws.read() {
Ok(Message::Text(text)) => ws.send(Message::Text(text)).unwrap_or_else(|e| panic!("server send: {e}")),
Ok(Message::Binary(bytes)) => ws.send(Message::Binary(bytes)).unwrap_or_else(|e| panic!("server send: {e}")),
Ok(Message::Close(_)) => {
let _ = ws.flush();
return;
}
Ok(_) => {}
Err(tungstenite::Error::Io(e)) if is_no_data(&e) => {}
Err(_) => return,
}
}
});
(port, handle)
}
#[test]
fn large_tls_messages_survive_read_timeouts_mid_record() {
let pki = pki();
let (port, server) = echo_server(&pki);
let handshake = handshake(
format!("wss://localhost:{port}/"),
Duration::from_secs(10),
Duration::from_millis(5),
Duration::from_millis(50),
Duration::from_secs(20),
);
let (commands, events, client) = start(handshake, Some(client_config(&pki)));
assert_eq!(events.recv_timeout(Duration::from_secs(10)).map(|(_, e)| e), Ok(WsLinkEvent::Opened));
for i in 0..12u8 {
let binary = (0..700 * 1024u32).map(|n| u8::try_from(n % 251).unwrap_or(0) ^ i).collect::<Vec<u8>>();
let text = format!("{i}:{}", "abcdefghij".repeat(30 * 1024));
for frame in [WsFrame::Binary(binary), WsFrame::Text(text)] {
commands.send(Command::Send(frame.clone())).unwrap_or_else(|e| panic!("{e}"));
match events.recv_timeout(Duration::from_secs(30)) {
Ok((_, WsLinkEvent::Frame(echo))) => assert!(echo == frame, "echo {i} differs"),
other => panic!("unexpected: {other:?}"),
}
}
}
commands.send(Command::Close(1000)).unwrap_or_else(|e| panic!("{e}"));
let last = client.join().unwrap_or_else(|_| panic!("client thread"));
assert!(matches!(last, WsLinkEvent::Closed { .. }), "{last:?}");
let _ = server.join();
}
fn trickling_server(prefix: &'static [u8]) -> u16 {
let listener = TcpListener::bind("127.0.0.1:0").unwrap_or_else(|e| panic!("{e}"));
let port = listener.local_addr().map(|a| a.port()).unwrap_or_else(|e| panic!("{e}"));
thread::spawn(move || {
let Ok((mut tcp, _)) = listener.accept() else { return };
let _ = tcp.set_read_timeout(Some(Duration::from_millis(200)));
let mut buf = [0u8; 4096];
let _ = tcp.read(&mut buf);
if tcp.write_all(prefix).is_err() {
return;
}
for _ in 0..20 {
thread::sleep(Duration::from_millis(300));
if tcp.write_all(b"x").is_err() {
return;
}
}
});
port
}
fn assert_connect_limit_holds(uri: String, tls: Option<Arc<rustls::ClientConfig>>) {
let started = Instant::now();
let handshake = handshake(uri, Duration::from_secs(1), Duration::from_millis(20), Duration::from_secs(15), Duration::from_secs(45));
let (_commands, events, client) = start(handshake, tls);
let last = client.join().unwrap_or_else(|_| panic!("client thread"));
assert!(matches!(&last, WsLinkEvent::Failed(BackendError::Timeout(why)) if why == "connect limit"), "{last:?}");
assert!(started.elapsed() < Duration::from_millis(2500), "the attempt took {:?}", started.elapsed());
assert!(events.try_recv().is_err(), "never reported open");
}
#[test]
fn a_trickled_http_handshake_stops_at_the_connect_limit() {
let port = trickling_server(b"HTTP/1.1 101 Switching Protocols\r\nX-Slow: ");
assert_connect_limit_holds(format!("ws://127.0.0.1:{port}/"), None);
}
#[test]
fn a_trickled_tls_handshake_stops_at_the_connect_limit() {
let port = trickling_server(&[0x16, 0x03, 0x03, 0x40, 0x00]);
let pki = pki();
assert_connect_limit_holds(format!("wss://localhost:{port}/"), Some(client_config(&pki)));
}
#[test]
fn a_trickled_incoming_frame_does_not_starve_sends_and_pings() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap_or_else(|e| panic!("{e}"));
let port = listener.local_addr().map(|a| a.port()).unwrap_or_else(|e| panic!("{e}"));
let received = Arc::new(AtomicUsize::new(0));
let during_trickle = Arc::new(AtomicUsize::new(0));
let (count, window) = (Arc::clone(&received), Arc::clone(&during_trickle));
thread::spawn(move || {
let Ok((tcp, _)) = listener.accept() else { return };
let Ok(ws) = tungstenite::accept(tcp.try_clone().unwrap_or_else(|e| panic!("{e}"))) else { return };
drop(ws);
let mut reader = tcp.try_clone().unwrap_or_else(|e| panic!("{e}"));
thread::spawn(move || {
let mut buf = [0u8; 4096];
while let Ok(n) = reader.read(&mut buf) {
if n == 0 {
break;
}
count.fetch_add(n, Ordering::SeqCst);
}
});
let mut writer = tcp;
let _ = writer.write_all(&[0x82, 126, 0xEA, 0x60]);
let before = received.load(Ordering::SeqCst);
let until = Instant::now() + Duration::from_secs(3);
while Instant::now() < until {
thread::sleep(Duration::from_millis(8));
if writer.write_all(b"z").is_err() {
return;
}
}
window.store(received.load(Ordering::SeqCst).saturating_sub(before), Ordering::SeqCst);
thread::sleep(Duration::from_secs(5));
});
let handshake = handshake(
format!("ws://127.0.0.1:{port}/"),
Duration::from_secs(5),
Duration::from_millis(20),
Duration::from_millis(50),
Duration::from_millis(500),
);
let (commands, events, client) = start(handshake, None);
assert_eq!(events.recv_timeout(Duration::from_secs(5)).map(|(_, e)| e), Ok(WsLinkEvent::Opened));
thread::sleep(Duration::from_millis(1000));
commands.send(Command::Send(WsFrame::Text("sent during the trickle".into()))).unwrap_or_else(|e| panic!("{e}"));
let last = client.join().unwrap_or_else(|_| panic!("client thread"));
assert!(during_trickle.load(Ordering::SeqCst) >= 200, "only {} bytes sent while the frame trickled", during_trickle.load(Ordering::SeqCst));
assert!(matches!(&last, WsLinkEvent::Failed(BackendError::Timeout(why)) if why.contains("dead")), "{last:?}");
}
#[test]
fn a_large_send_to_a_slow_reader_does_not_kill_a_live_connection() {
let listener = TcpListener::bind("127.0.0.1:0").unwrap_or_else(|e| panic!("{e}"));
let port = listener.local_addr().map(|a| a.port()).unwrap_or_else(|e| panic!("{e}"));
let received = Arc::new(AtomicUsize::new(0));
let count = Arc::clone(&received);
thread::spawn(move || {
let Ok((tcp, _)) = listener.accept() else { return };
let Ok(ws) = tungstenite::accept(tcp.try_clone().unwrap_or_else(|e| panic!("{e}"))) else { return };
drop(ws);
let mut writer = tcp.try_clone().unwrap_or_else(|e| panic!("{e}"));
thread::spawn(move || {
for _ in 0..60 {
if writer.write_all(&[0x81, 5, b'a', b'l', b'i', b'v', b'e']).is_err() {
return;
}
thread::sleep(Duration::from_millis(100));
}
});
let mut reader = tcp;
thread::sleep(Duration::from_millis(2500));
let mut buf = vec![0u8; 1 << 16];
let until = Instant::now() + Duration::from_secs(20);
while Instant::now() < until {
match reader.read(&mut buf) {
Ok(0) | Err(_) => break,
Ok(n) => {
count.fetch_add(n, Ordering::SeqCst);
}
}
}
});
let handshake =
handshake(format!("ws://127.0.0.1:{port}/"), Duration::from_secs(5), Duration::from_millis(20), Duration::from_millis(200), Duration::from_secs(1));
let (commands, events, client) = start(handshake, None);
assert_eq!(events.recv_timeout(Duration::from_secs(5)).map(|(_, e)| e), Ok(WsLinkEvent::Opened));
commands.send(Command::Send(WsFrame::Binary(vec![7u8; 32 << 20]))).unwrap_or_else(|e| panic!("{e}"));
let until = Instant::now() + Duration::from_secs(5);
let mut frames = 0;
while Instant::now() < until {
match events.recv_timeout(Duration::from_millis(100)) {
Ok((_, WsLinkEvent::Frame(_))) => frames += 1,
Ok((_, other)) => panic!("the live connection ended: {other:?}"),
Err(_) => {}
}
}
assert!(frames >= 20, "only {frames} frames");
assert!(received.load(Ordering::SeqCst) >= 32 << 20, "the server got {} bytes", received.load(Ordering::SeqCst));
commands.send(Command::Close(1000)).unwrap_or_else(|e| panic!("{e}"));
let _ = client.join();
}
#[test]
fn huge_durations_never_overflow() {
let far = deadline_after(Duration::MAX);
assert!(far > Instant::now());
let listener = TcpListener::bind("127.0.0.1:0").unwrap_or_else(|e| panic!("{e}"));
let tcp = TcpStream::connect(listener.local_addr().unwrap_or_else(|e| panic!("{e}"))).unwrap_or_else(|e| panic!("{e}"));
let mut timed = TimedTcp::new(tcp, Some(Instant::now()), Duration::MAX);
let mut buf = [0u8; 4];
assert_eq!(timed.read(&mut buf).map_err(|e| e.kind()), Err(io::ErrorKind::TimedOut));
}
}