diff --git a/moli-curl/src/websocket.rs b/moli-curl/src/websocket.rs index 39bad96ab8..ba13fe4cba 100644 --- a/moli-curl/src/websocket.rs +++ b/moli-curl/src/websocket.rs @@ -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 { diff --git a/moli-curl/src/websocket/owner.rs b/moli-curl/src/websocket/owner.rs index 97c541c23c..249fb54a96 100644 --- a/moli-curl/src/websocket/owner.rs +++ b/moli-curl/src/websocket/owner.rs @@ -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, dns: CurlDnsOwnerResidence, + 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()); } } diff --git a/moli-curl/src/websocket/scheduling.rs b/moli-curl/src/websocket/scheduling.rs new file mode 100644 index 0000000000..6a3eccaf75 --- /dev/null +++ b/moli-curl/src/websocket/scheduling.rs @@ -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, + fds: Vec, +} + +impl SocketPoll { + pub(super) fn wait( + &mut self, + multi: &Multi, + sessions: &mut HashMap, + 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(()) + } +} diff --git a/moli-curl/src/websocket/session.rs b/moli-curl/src/websocket/session.rs index 812ebc0d96..6a8d4add6b 100644 --- a/moli-curl/src/websocket/session.rs +++ b/moli-curl/src/websocket/session.rs @@ -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, 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 { // 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 { - 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); diff --git a/moli-curl/src/websocket/tests.rs b/moli-curl/src/websocket/tests.rs index dc38a0ecbe..8bebdd9ded 100644 --- a/moli-curl/src/websocket/tests.rs +++ b/moli-curl/src/websocket/tests.rs @@ -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<()>) { diff --git a/moli-curl/src/websocket/tests/readiness.rs b/moli-curl/src/websocket/tests/readiness.rs new file mode 100644 index 0000000000..0d255468ae --- /dev/null +++ b/moli-curl/src/websocket/tests/readiness.rs @@ -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(); +}