fix(websocket): schedule native I/O by readiness

This commit is contained in:
ldm0
2026-09-11 01:47:10 +08:00
committed by Donough Liu
parent cd3a2b1679
commit 14fed394be
6 changed files with 392 additions and 24 deletions
+15
View File
@@ -3,9 +3,16 @@
//! This layer transports frames. Browser handshake policy, message assembly and
//! close-handshake semantics belong to the caller. Dropping the receiver cancels
//! its session, including DNS and handshake work, independently of queue capacity.
//!
//! Internally, owner coordinates DNS and scheduling; session owns the native
//! handle and its Opening/Open/ReceivedClose lifecycle; scheduling holds I/O
//! admission state and maps polled sockets back to their sessions. Returning
//! AGAIN parks that I/O until its socket is signalled. Application wakeups only
//! resume paused work, such as a new frame or restored receive capacity.
mod owner;
mod request;
mod scheduling;
mod session;
#[cfg(test)]
mod tests;
@@ -124,6 +131,10 @@ struct Control {
read_blocked: tokio::sync::Notify,
#[cfg(test)]
write_blocked: tokio::sync::Notify,
#[cfg(test)]
read_attempts: std::sync::atomic::AtomicUsize,
#[cfg(test)]
read_waiting: tokio::sync::Notify,
}
impl Control {
@@ -329,6 +340,10 @@ impl CurlWebSocketRuntime {
read_blocked: tokio::sync::Notify::new(),
#[cfg(test)]
write_blocked: tokio::sync::Notify::new(),
#[cfg(test)]
read_attempts: std::sync::atomic::AtomicUsize::new(0),
#[cfg(test)]
read_waiting: tokio::sync::Notify::new(),
});
let (event_tx, events) = mpsc::channel(MAX_PENDING_EVENTS);
let io = SessionIo {
+20 -17
View File
@@ -18,6 +18,7 @@ use curl::{
use super::{
SessionIo, Submission,
request::{self, Handshake},
scheduling::SocketPoll,
session::{Session, Step},
};
use crate::{CurlDnsResolution, CurlTransferId, dns_adapter::CurlDnsOwnerResidence};
@@ -36,6 +37,7 @@ struct Owner {
multi: Multi,
sessions: HashMap<CurlTransferId, Session>,
dns: CurlDnsOwnerResidence<CurlTransferId, Pending>,
poll: SocketPoll,
}
pub(super) fn run(
@@ -47,15 +49,15 @@ pub(super) fn run(
multi: Multi::new(),
sessions: HashMap::new(),
dns: CurlDnsOwnerResidence::default(),
poll: SocketPoll::default(),
};
let _ = waker_tx.send(owner.multi.waker());
while !shutdown.load(Ordering::Acquire) {
owner.accept_submissions(&submissions);
owner.resolve_dns();
owner.advance_handshakes();
if !owner.advance_sessions() {
owner.wait_for_work();
}
let progressed = owner.advance_sessions();
owner.wait_for_work(progressed);
}
owner.fail_sessions("curl WebSocket runtime shut down");
for pending in owner.dns.drain() {
@@ -190,26 +192,27 @@ impl Owner {
}
}
fn wait_for_work(&mut self) {
let mut fds: Vec<_> = self
.sessions
.values()
.filter_map(Session::wait_fd)
.collect();
fn wait_for_work(&mut self, progressed: bool) {
let deadline = self
.sessions
.values()
.filter_map(Session::handshake_deadline)
.chain(self.dns.next_deadline(|pending| Some(pending.deadline)))
.min();
let timeout = deadline
.map(|deadline| {
deadline
.saturating_duration_since(Instant::now())
.min(IDLE_WAIT)
})
.unwrap_or(IDLE_WAIT);
if let Err(error) = self.multi.poll(&mut fds, timeout) {
let timeout = if progressed {
// Keep draining local work, but collect other sockets' readiness on
// every turn so a busy session cannot starve a newly readable one.
Duration::ZERO
} else {
deadline
.map(|deadline| {
deadline
.saturating_duration_since(Instant::now())
.min(IDLE_WAIT)
})
.unwrap_or(IDLE_WAIT)
};
if let Err(error) = self.poll.wait(&self.multi, &mut self.sessions, timeout) {
self.fail_sessions(&error.to_string());
}
}
+86
View File
@@ -0,0 +1,86 @@
//! Admission and socket readiness are separate. Paused I/O needs application
//! work/capacity; WaitingForSocket I/O has already returned AGAIN. A successful
//! operation stays Runnable so libcurl's buffered data can be drained as well.
use std::{collections::HashMap, time::Duration};
use curl::multi::{Multi, WaitFd};
use super::session::Session;
use crate::CurlTransferId;
#[derive(Default, PartialEq, Eq)]
pub(super) enum IoState {
#[default]
Paused,
Runnable,
WaitingForSocket,
}
impl IoState {
pub(super) fn can_run(&mut self, enabled: bool) -> bool {
if !enabled {
self.pause();
} else if *self == Self::Paused {
*self = Self::Runnable;
}
*self == Self::Runnable
}
pub(super) fn pause(&mut self) {
*self = Self::Paused;
}
pub(super) fn would_block(&mut self) {
*self = Self::WaitingForSocket;
}
pub(super) fn socket_ready(&mut self) {
if self.waiting_for_socket() {
*self = Self::Runnable;
}
}
pub(super) fn waiting_for_socket(&self) -> bool {
*self == Self::WaitingForSocket
}
}
/// Keeps each extra fd paired with its session across curl_multi_poll.
/// libcurl also polls opening handshakes and the cross-thread waker here.
#[derive(Default)]
pub(super) struct SocketPoll {
ids: Vec<CurlTransferId>,
fds: Vec<WaitFd>,
}
impl SocketPoll {
pub(super) fn wait(
&mut self,
multi: &Multi,
sessions: &mut HashMap<CurlTransferId, Session>,
timeout: Duration,
) -> Result<(), curl::MultiError> {
self.ids.clear();
self.fds.clear();
for (id, session) in sessions.iter() {
if let Some(fd) = session.wait_fd() {
self.ids.push(*id);
self.fds.push(fd);
}
}
multi.poll(&mut self.fds, timeout)?;
for (id, fd) in self.ids.iter().zip(&self.fds) {
if fd.received_read() || fd.received_write() {
// A socket event permits a retry, not guaranteed progress.
// curl maps HUP/ERR into read/write bits; retry both directions
// so a write-only session can also observe peer shutdown.
sessions
.get_mut(id)
.expect("polled session is resident")
.socket_ready();
}
}
Ok(())
}
}
+37 -7
View File
@@ -8,7 +8,7 @@ use curl::{
multi::{Easy2Handle, Multi, WaitFd},
};
use super::{CurlWebSocketEvent, SessionIo, WsFlags, request::Handshake};
use super::{CurlWebSocketEvent, SessionIo, WsFlags, request::Handshake, scheduling::IoState};
use crate::CurlTransferId;
const CHUNK_BYTES: usize = 16 * 1024;
@@ -32,6 +32,8 @@ pub(super) struct Session {
handle: Easy2Handle<Handshake>,
io: SessionIo,
phase: Phase,
reading: IoState,
writing: IoState,
}
impl Session {
@@ -66,6 +68,8 @@ impl Session {
handle,
io,
phase: Phase::Opening { deadline },
reading: IoState::default(),
writing: IoState::default(),
})
}
Err(error) => {
@@ -143,9 +147,12 @@ impl Session {
fn write_pending(&mut self) -> Result<bool, String> {
// Keep the single frame resident across partial nonblocking writes.
let mut pending = self.io.control.send.lock();
let Some(send) = &mut *pending else {
if !self.writing.can_run(pending.is_some()) {
return Ok(false);
};
}
let send = pending
.as_mut()
.expect("runnable write has a pending frame");
match self
.handle
.ws_send(&send.frame.data[send.offset..], 0, send.frame.flags)
@@ -154,11 +161,13 @@ impl Session {
send.offset += count;
if send.offset == send.frame.data.len() {
let send = pending.take().expect("completed frame");
self.writing.pause();
let _ = send.completed.send(send.offset);
}
Ok(true)
}
Err(error) if error.is_again() => {
self.writing.would_block();
#[cfg(test)]
self.io.control.write_blocked.notify_one();
Ok(false)
@@ -172,12 +181,13 @@ impl Session {
}
fn read_chunk(&mut self) -> Result<Step, String> {
if !self.reading_enabled() {
if !self.reading.can_run(self.reading_enabled()) {
return Ok(Step::Idle);
}
let permit = match self.io.events.try_reserve() {
Ok(permit) => permit,
Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => {
self.reading.pause();
#[cfg(test)]
self.io.control.read_blocked.notify_one();
return Ok(Step::Idle);
@@ -185,6 +195,11 @@ impl Session {
Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => return Ok(Step::Closed),
};
let mut data = vec![0; CHUNK_BYTES];
#[cfg(test)]
self.io
.control
.read_attempts
.fetch_add(1, Ordering::Relaxed);
match self.handle.ws_recv(&mut data) {
Ok((count, frame)) => {
data.truncate(count);
@@ -196,7 +211,12 @@ impl Session {
permit.send(CurlWebSocketEvent::Chunk { data, frame });
Ok(Step::Progress)
}
Err(error) if error.is_again() => Ok(Step::Idle),
Err(error) if error.is_again() => {
self.reading.would_block();
#[cfg(test)]
self.io.control.read_waiting.notify_one();
Ok(Step::Idle)
}
Err(error) if error.is_got_nothing() => Ok(Step::Closed),
Err(error) => Err(format!("WebSocket receive failed: {error}")),
}
@@ -206,17 +226,27 @@ impl Session {
if self.handshake_deadline().is_some() {
return None;
}
let reading = self.reading_enabled() && self.io.events.capacity() > 0;
let writing = self.io.control.send.lock().is_some();
let reading = self.reading.waiting_for_socket()
&& self.reading_enabled()
&& self.io.events.capacity() > 0;
let writing = self.writing.waiting_for_socket();
if !reading && !writing {
return None;
}
let mut fd = WaitFd::new();
fd.set_fd(self.handle.active_socket().ok()??);
// curl-rust exposes AGAIN without a TLS wait direction. These are the
// operation's read/write interests; cross-direction TLS waits would
// require extending the native contract here.
fd.poll_on_read(reading).poll_on_write(writing);
Some(fd)
}
pub(super) fn socket_ready(&mut self) {
self.reading.socket_ready();
self.writing.socket_ready();
}
pub(super) fn finish(self, multi: &mut Multi, result: Result<(), String>) {
// Physically release the socket before publishing terminal channel closure.
let _ = multi.remove2(self.handle);
+2
View File
@@ -9,6 +9,8 @@ use tokio_tungstenite::tungstenite::{self, Message, handshake::derive_accept_key
use super::*;
mod readiness;
const DEADLINE: Duration = Duration::from_secs(10);
fn server(handler: impl FnOnce(TcpStream) + Send + 'static) -> (String, thread::JoinHandle<()>) {
+232
View File
@@ -0,0 +1,232 @@
use super::*;
#[tokio::test]
async fn native_active_sessions_do_not_probe_idle_sockets() {
let (idle_url, idle_task) = server(|stream| {
let mut socket = tungstenite::accept(stream).unwrap();
assert_eq!(socket.read().unwrap(), Message::Text("wake".into()));
socket.send(Message::Text("awake".into())).unwrap();
assert!(socket.read().is_err());
});
let runtime = CurlWebSocketRuntime::new().unwrap();
let mut idle = runtime
.connect(CurlWebSocketRequest::new(idle_url))
.unwrap();
opened(&mut idle).await;
let idle_sender = idle.sender();
idle_sender.set_reading(true);
timeout(DEADLINE, idle_sender.control.read_waiting.notified())
.await
.unwrap();
let before = idle_sender.control.read_attempts.load(Ordering::Acquire);
let (active_url, active_task) = server(|mut stream| {
let tail: Vec<_> = (0..128).flat_map(|i| [0x82, 1, i]).collect();
upgrade(&mut stream, &tail);
assert_eq!(stream.read(&mut [0]).unwrap(), 0);
});
let mut active = runtime
.connect(CurlWebSocketRequest::new(active_url))
.unwrap();
opened(&mut active).await;
active.sender().set_reading(true);
for expected in 0..128 {
assert!(matches!(event(&mut active).await,
CurlWebSocketEvent::Chunk { data, .. } if data == [expected]));
}
assert_eq!(
idle_sender.control.read_attempts.load(Ordering::Acquire),
before,
"another connection's progress and delivery wakes must not reprobe an idle socket"
);
// A new send wakes this idle session independently of its waiting reader.
assert_eq!(
timeout(
DEADLINE,
idle_sender.send_frame(CurlWebSocketSend {
flags: WsFlags::TEXT,
data: b"wake".to_vec(),
})
)
.await
.unwrap()
.unwrap(),
4
);
assert!(matches!(event(&mut idle).await,
CurlWebSocketEvent::Chunk { data, .. } if data == b"awake"));
drop(idle);
drop(active);
idle_task.join().unwrap();
active_task.join().unwrap();
}
#[tokio::test]
async fn native_socket_readiness_is_serviced_during_continuous_traffic() {
let (url, task) = server(|stream| {
let mut socket = tungstenite::accept(stream).unwrap();
while socket
.send(Message::Binary(vec![3; MAX_SEND_FRAME_BYTES].into()))
.is_ok()
{}
});
let (signal_tx, signal_rx) = std::sync::mpsc::channel();
let (quiet_url, quiet_task) = server(move |mut stream| {
upgrade(&mut stream, &[]);
signal_rx.recv_timeout(DEADLINE).unwrap();
stream.write_all(&[0x81, 2, b'o', b'k']).unwrap();
assert_eq!(stream.read(&mut [0]).unwrap(), 0);
});
let runtime = CurlWebSocketRuntime::new().unwrap();
let mut quiet = runtime
.connect(CurlWebSocketRequest::new(quiet_url))
.unwrap();
opened(&mut quiet).await;
quiet.sender().set_reading(true);
timeout(DEADLINE, quiet.sender.control.read_waiting.notified())
.await
.unwrap();
let mut active = runtime.connect(CurlWebSocketRequest::new(url)).unwrap();
opened(&mut active).await;
active.sender().set_reading(true);
let (started_tx, started_rx) = oneshot::channel();
let (stop_tx, stop_rx) = oneshot::channel();
let drain = tokio::spawn(async move {
let mut started = Some(started_tx);
tokio::pin!(stop_rx);
loop {
tokio::select! {
_ = &mut stop_rx => break,
next = active.recv() => {
assert!(matches!(next, Some(CurlWebSocketEvent::Chunk { data, .. })
if !data.is_empty() && data.iter().all(|byte| *byte == 3)));
if let Some(started) = started.take() { started.send(()).unwrap(); }
}
}
}
});
timeout(DEADLINE, started_rx).await.unwrap().unwrap();
signal_tx.send(()).unwrap();
assert!(matches!(event(&mut quiet).await,
CurlWebSocketEvent::Chunk { data, .. } if data == b"ok"));
stop_tx.send(()).unwrap();
timeout(DEADLINE, drain).await.unwrap().unwrap();
drop(quiet);
task.join().unwrap();
quiet_task.join().unwrap();
}
#[tokio::test]
async fn native_wss_buffered_frames_resume_after_backpressure() {
use rustls::{
ServerConfig, ServerConnection, StreamOwned,
pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer},
};
let certificate = rcgen::generate_simple_self_signed(vec!["localhost".to_owned()]).unwrap();
let config = Arc::new(
ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(
vec![CertificateDer::from(certificate.cert.der().to_vec())],
PrivateKeyDer::from(PrivatePkcs8KeyDer::from(
certificate.key_pair.serialize_der(),
)),
)
.unwrap(),
);
let (url, task) = server(move |stream| {
let tls = StreamOwned::new(ServerConnection::new(config).unwrap(), stream);
let mut socket = tungstenite::accept(tls).unwrap();
// One TLS write coalesces enough frames to exceed the delivery queue.
let wire: Vec<_> = (0..32).flat_map(|i| [0x82, 1, i]).collect();
socket.get_mut().write_all(&wire).unwrap();
socket.get_mut().flush().unwrap();
assert!(socket.read().is_err());
});
let runtime = CurlWebSocketRuntime::new().unwrap();
let mut request = CurlWebSocketRequest::new(url.replacen("ws://", "wss://", 1));
// This fixture tests TLS buffering; certificate policy has separate coverage.
request.tls.verify = false;
let mut connection = runtime.connect(request).unwrap();
opened(&mut connection).await;
let sender = connection.sender();
sender.set_reading(true);
timeout(DEADLINE, sender.control.read_blocked.notified())
.await
.unwrap();
sender.set_reading(false);
for expected in 0..MAX_PENDING_EVENTS {
assert!(matches!(event(&mut connection).await,
CurlWebSocketEvent::Chunk { data, .. } if data == [expected as u8]));
}
// No further server writes: resume must also drain libcurl/TLS buffers.
sender.set_reading(true);
for expected in MAX_PENDING_EVENTS..32 {
assert!(matches!(event(&mut connection).await,
CurlWebSocketEvent::Chunk { data, .. } if data == [expected as u8]));
}
drop(connection);
task.join().unwrap();
}
#[tokio::test]
async fn native_peer_eof_wakes_waiting_reader() {
let (finish_tx, finish_rx) = std::sync::mpsc::channel();
let (url, task) = server(move |mut stream| {
upgrade(&mut stream, &[]);
finish_rx.recv_timeout(DEADLINE).unwrap();
});
let runtime = CurlWebSocketRuntime::new().unwrap();
let mut connection = runtime.connect(CurlWebSocketRequest::new(url)).unwrap();
opened(&mut connection).await;
connection.sender().set_reading(true);
timeout(DEADLINE, connection.sender.control.read_waiting.notified())
.await
.unwrap();
finish_tx.send(()).unwrap();
assert!(matches!(
event(&mut connection).await,
CurlWebSocketEvent::Closed { result: Ok(()) }
));
assert!(connection.recv().await.is_none());
task.join().unwrap();
}
#[tokio::test]
async fn native_peer_shutdown_settles_blocked_writer_with_reads_paused() {
let (finish_tx, finish_rx) = std::sync::mpsc::channel();
let (url, task) = server(move |stream| {
let _socket = tungstenite::accept(stream).unwrap();
finish_rx.recv_timeout(DEADLINE).unwrap();
});
let runtime = CurlWebSocketRuntime::new().unwrap();
let mut connection = runtime.connect(CurlWebSocketRequest::new(url)).unwrap();
opened(&mut connection).await;
let sender = connection.sender();
let send_until_blocked = async {
for _ in 0..128 {
let send = sender.send_frame(CurlWebSocketSend {
flags: WsFlags::BINARY,
data: vec![9; MAX_SEND_FRAME_BYTES],
});
tokio::pin!(send);
tokio::select! {
biased;
_ = sender.control.write_blocked.notified() => {
finish_tx.send(()).unwrap();
assert!(send.await.is_err());
return;
}
result = &mut send => { assert_eq!(result.unwrap(), MAX_SEND_FRAME_BYTES); }
}
}
panic!("server backpressure must reach the writer");
};
timeout(DEADLINE, send_until_blocked).await.unwrap();
assert!(matches!(
event(&mut connection).await,
CurlWebSocketEvent::Closed { result: Err(_) }
));
task.join().unwrap();
}