From cd3a2b167930f2ae91245af0fc3327cf0ebffc40 Mon Sep 17 00:00:00 2001 From: ldm0 Date: Thu, 10 Sep 2026 14:49:58 +0800 Subject: [PATCH] refactor(websocket): encapsulate native session lifecycle --- moli-curl/src/websocket.rs | 1 + moli-curl/src/websocket/owner.rs | 310 ++++++++++------------------- moli-curl/src/websocket/session.rs | 225 +++++++++++++++++++++ 3 files changed, 328 insertions(+), 208 deletions(-) create mode 100644 moli-curl/src/websocket/session.rs diff --git a/moli-curl/src/websocket.rs b/moli-curl/src/websocket.rs index 5b7c876d80..39bad96ab8 100644 --- a/moli-curl/src/websocket.rs +++ b/moli-curl/src/websocket.rs @@ -6,6 +6,7 @@ mod owner; mod request; +mod session; #[cfg(test)] mod tests; diff --git a/moli-curl/src/websocket/owner.rs b/moli-curl/src/websocket/owner.rs index b41f2bb640..97c541c23c 100644 --- a/moli-curl/src/websocket/owner.rs +++ b/moli-curl/src/websocket/owner.rs @@ -1,3 +1,6 @@ +//! Coordinates admission, DNS, handshakes and scheduling. Native frame I/O and +//! the lifetime of an attached easy handle belong to Session. + use std::{ collections::HashMap, sync::{ @@ -9,17 +12,16 @@ use std::{ use curl::{ easy::Easy2, - multi::{Easy2Handle, Multi, MultiWaker, WaitFd}, + multi::{Multi, MultiWaker}, }; use super::{ - CurlWebSocketEvent, SessionIo, Submission, + SessionIo, Submission, request::{self, Handshake}, + session::{Session, Step}, }; use crate::{CurlDnsResolution, CurlTransferId, dns_adapter::CurlDnsOwnerResidence}; -const CHUNK_BYTES: usize = 16 * 1024; -const IO_BUDGET: usize = 8; const IDLE_WAIT: Duration = Duration::from_secs(1); struct Pending { @@ -30,12 +32,10 @@ struct Pending { io: SessionIo, } -struct Session { - handle: Easy2Handle, - deadline: Instant, - open: bool, - received_close: bool, - io: SessionIo, +struct Owner { + multi: Multi, + sessions: HashMap, + dns: CurlDnsOwnerResidence, } pub(super) fn run( @@ -43,11 +43,30 @@ pub(super) fn run( waker_tx: crossbeam_channel::Sender, shutdown: Arc, ) { - let mut multi = Multi::new(); - let _ = waker_tx.send(multi.waker()); - let mut sessions = HashMap::::new(); - let mut dns = CurlDnsOwnerResidence::::default(); + let mut owner = Owner { + multi: Multi::new(), + sessions: HashMap::new(), + dns: CurlDnsOwnerResidence::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(); + } + } + owner.fail_sessions("curl WebSocket runtime shut down"); + for pending in owner.dns.drain() { + pending + .io + .finish(Err("curl WebSocket runtime shut down".to_owned())); + } +} + +impl Owner { + fn accept_submissions(&mut self, submissions: &crossbeam_channel::Receiver) { for submission in submissions.try_iter().take(super::SESSION_CAPACITY) { if submission.io.cancelled() { continue; @@ -74,18 +93,24 @@ pub(super) fn run( io: submission.io, }; match pending.dns.target().cloned() { - Some(target) => dns.start(pending.id, pending, target, multi.waker()), - None => start(&mut multi, &mut sessions, pending), + Some(target) => self + .dns + .start(pending.id, pending, target, self.multi.waker()), + None => self.start(pending), } } - for pending in dns + } + + fn resolve_dns(&mut self) { + for pending in self + .dns .take_matching(|pending| pending.io.cancelled() || pending.deadline <= Instant::now()) { pending .io .finish(Err("WebSocket DNS cancelled or timed out".to_owned())); } - while let Some(ready) = dns.try_claim_next() { + while let Some(ready) = self.dns.try_claim_next() { let mut pending = ready.pending; if pending.io.cancelled() { continue; @@ -100,17 +125,30 @@ pub(super) fn run( .map_err(|error| error.to_string()) }); match result { - Ok(()) => start(&mut multi, &mut sessions, pending), + Ok(()) => self.start(pending), Err(error) => pending.io.finish(Err(error)), } } - if let Err(error) = multi.perform() { - for (_, session) in sessions.drain() { - finish(&mut multi, session, Err(error.to_string())); - } + } + + fn start(&mut self, pending: Pending) { + if let Some(session) = Session::attach( + &mut self.multi, + pending.id, + pending.easy, + pending.io, + pending.deadline, + ) { + self.sessions.insert(pending.id, session); + } + } + + fn advance_handshakes(&mut self) { + if let Err(error) = self.multi.perform() { + self.fail_sessions(&error.to_string()); } let mut completed = Vec::new(); - multi.messages(|message| { + self.multi.messages(|message| { if let (Ok(token), Some(result)) = (message.token(), message.result()) && let Some(id) = CurlTransferId::from_token(token) { @@ -118,60 +156,51 @@ pub(super) fn run( } }); for (id, result) in completed { - let Some(session) = sessions.get_mut(&id) else { - continue; - }; - if session.open { - continue; - } - let handshake = session.handle.get_mut(); - let native_result = match handshake.error.take() { - Some(error) => Err(error), - None => result.map_err(|error| error.to_string()), - }; - let event = CurlWebSocketEvent::Handshake { - request: std::mem::take(&mut handshake.request), - response: std::mem::take(&mut handshake.response), - result: native_result.clone(), - }; - // No data events precede the handshake, so this queue has a free slot. - if session.io.events.try_send(event).is_err() || native_result.is_err() { - let session = sessions.remove(&id).expect("completed session is resident"); - finish(&mut multi, session, native_result); - } else { - session.open = true; + if let Some(session) = self.sessions.get_mut(&id) { + let step = session.complete_handshake(result); + self.apply_step(id, step); } } + } + fn advance_sessions(&mut self) -> bool { let mut retired = Vec::new(); let mut progressed = false; - for (id, session) in &mut sessions { - if session.io.cancelled() { - retired.push((*id, Ok(()))); - } else if !session.open && session.deadline <= Instant::now() { - retired.push((*id, Err("WebSocket handshake timed out".to_owned()))); - } else if session.open { - match session.drive() { - Ok(progress) => progressed |= progress, - Err(result) => retired.push((*id, result)), - } + for (id, session) in &mut self.sessions { + match session.advance() { + Ok(Step::Progress) => progressed = true, + Ok(Step::Idle) => {} + terminal => retired.push((*id, terminal)), } } - for (id, result) in retired { - if let Some(session) = sessions.remove(&id) { - finish(&mut multi, session, result); - } - } - if progressed { - continue; + for (id, step) in retired { + self.apply_step(id, step); } + progressed + } - let mut fds: Vec<_> = sessions.values().filter_map(Session::wait_fd).collect(); - let deadline = sessions + fn apply_step(&mut self, id: CurlTransferId, step: Result) { + let result = match step { + Ok(Step::Progress | Step::Idle) => return, + Ok(Step::Closed) => Ok(()), + Err(error) => Err(error), + }; + if let Some(session) = self.sessions.remove(&id) { + session.finish(&mut self.multi, result); + } + } + + fn wait_for_work(&mut self) { + let mut fds: Vec<_> = self + .sessions .values() - .filter(|session| !session.open) - .map(|session| session.deadline) - .chain(dns.next_deadline(|pending| Some(pending.deadline))) + .filter_map(Session::wait_fd) + .collect(); + 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| { @@ -180,149 +209,14 @@ pub(super) fn run( .min(IDLE_WAIT) }) .unwrap_or(IDLE_WAIT); - // poll includes libcurl's handshake sockets plus our upgraded sockets. - // Completed CONNECT_ONLY handles deliberately remain attached to this Multi. - if let Err(error) = multi.poll(&mut fds, timeout) { - for (_, session) in sessions.drain() { - finish(&mut multi, session, Err(error.to_string())); - } + if let Err(error) = self.multi.poll(&mut fds, timeout) { + self.fail_sessions(&error.to_string()); } } - for (_, session) in sessions.drain() { - finish( - &mut multi, - session, - Err("curl WebSocket runtime shut down".to_owned()), - ); - } - for pending in dns.drain() { - pending - .io - .finish(Err("curl WebSocket runtime shut down".to_owned())); - } -} - -fn start(multi: &mut Multi, sessions: &mut HashMap, mut pending: Pending) { - if pending.io.cancelled() { - return; - } - let remaining = pending.deadline.saturating_duration_since(Instant::now()); - if remaining.is_zero() { - pending - .io - .finish(Err("WebSocket handshake timed out".to_owned())); - return; - } - // timeout belongs to the transfer (opening handshake), not the upgraded session. - if let Err(error) = pending.easy.timeout(remaining) { - pending.io.finish(Err(error.to_string())); - return; - } - match multi.add2(pending.easy) { - Ok(mut handle) => { - if let Err(error) = handle.set_token(pending.id.token()) { - let _ = multi.remove2(handle); - pending.io.finish(Err(error.to_string())); - return; - } - sessions.insert( - pending.id, - Session { - handle, - deadline: pending.deadline, - open: false, - received_close: false, - io: pending.io, - }, - ); - } - Err(error) => pending.io.finish(Err(error.to_string())), - } -} - -fn finish(multi: &mut Multi, session: Session, result: std::result::Result<(), String>) { - // Physically release the socket before publishing terminal channel closure. - let _ = multi.remove2(session.handle); - session.io.finish(result); -} - -impl Session { - /// Err carries the terminal transport result; Ok reports useful work only. - fn drive(&mut self) -> std::result::Result> { - let mut progressed = false; - for _ in 0..IO_BUDGET { - if self.io.cancelled() { - return Err(Ok(())); - } - { - // Keep the single frame resident across partial nonblocking writes. - let mut pending = self.io.control.send.lock(); - if let Some(send) = &mut *pending { - match self - .handle - .ws_send(&send.frame.data[send.offset..], 0, send.frame.flags) - { - Ok(count) => { - send.offset += count; - progressed = true; - if send.offset == send.frame.data.len() { - let send = pending.take().expect("completed frame"); - let _ = send.completed.send(send.offset); - } - } - Err(error) if error.is_again() => { - #[cfg(test)] - self.io.control.write_blocked.notify_one(); - } - Err(error) => return Err(Err(format!("WebSocket send failed: {error}"))), - } - } - } - if self.received_close || !self.io.control.reading.load(Ordering::Acquire) { - break; - } - let permit = match self.io.events.try_reserve() { - Ok(permit) => permit, - Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => { - #[cfg(test)] - self.io.control.read_blocked.notify_one(); - break; - } - Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => return Err(Ok(())), - }; - let mut data = vec![0; CHUNK_BYTES]; - match self.handle.ws_recv(&mut data) { - Ok((count, frame)) => { - data.truncate(count); - // Let the caller answer Close before probing EOF. A peer can - // half-close TCP in the same packet as its Close frame. - self.received_close = - frame.flags().contains(super::WsFlags::CLOSE) && frame.bytes_left() == 0; - permit.send(CurlWebSocketEvent::Chunk { data, frame }); - progressed = true; - } - Err(error) if error.is_again() => break, - Err(error) if error.is_got_nothing() => return Err(Ok(())), - Err(error) => return Err(Err(format!("WebSocket receive failed: {error}"))), - } - } - Ok(progressed) - } - fn wait_fd(&self) -> Option { - if !self.open { - return None; + fn fail_sessions(&mut self, error: &str) { + for (_, session) in self.sessions.drain() { + session.finish(&mut self.multi, Err(error.to_owned())); } - let reading = self.io.events.capacity() > 0 - && !self.received_close - && self.io.control.reading.load(Ordering::Acquire); - let writing = self.io.control.send.lock().is_some(); - if !reading && !writing { - return None; - } - let mut fd = WaitFd::new(); - fd.set_fd(self.handle.active_socket().ok()??); - fd.poll_on_read(reading).poll_on_write(writing); - Some(fd) } } diff --git a/moli-curl/src/websocket/session.rs b/moli-curl/src/websocket/session.rs new file mode 100644 index 0000000000..812ebc0d96 --- /dev/null +++ b/moli-curl/src/websocket/session.rs @@ -0,0 +1,225 @@ +//! One attached libcurl handle. A completed CONNECT_ONLY transfer becomes an +//! open session; it stays in Multi until finish releases the socket. + +use std::{sync::atomic::Ordering, time::Instant}; + +use curl::{ + easy::Easy2, + multi::{Easy2Handle, Multi, WaitFd}, +}; + +use super::{CurlWebSocketEvent, SessionIo, WsFlags, request::Handshake}; +use crate::CurlTransferId; + +const CHUNK_BYTES: usize = 16 * 1024; +const IO_BUDGET: usize = 8; + +/// Result of advancing one native operation. Failures use Result::Err. +pub(super) enum Step { + Progress, + Idle, + /// EOF, cancellation or a dropped receiver. Browser close policy is external. + Closed, +} + +enum Phase { + Opening { deadline: Instant }, + Open, + ReceivedClose, +} + +pub(super) struct Session { + handle: Easy2Handle, + io: SessionIo, + phase: Phase, +} + +impl Session { + pub(super) fn attach( + multi: &mut Multi, + id: CurlTransferId, + mut easy: Easy2, + io: SessionIo, + deadline: Instant, + ) -> Option { + if io.cancelled() { + return None; + } + let remaining = deadline.saturating_duration_since(Instant::now()); + if remaining.is_zero() { + io.finish(Err("WebSocket handshake timed out".to_owned())); + return None; + } + // timeout belongs to the opening handshake, not the upgraded session. + if let Err(error) = easy.timeout(remaining) { + io.finish(Err(error.to_string())); + return None; + } + match multi.add2(easy) { + Ok(mut handle) => { + if let Err(error) = handle.set_token(id.token()) { + let _ = multi.remove2(handle); + io.finish(Err(error.to_string())); + return None; + } + Some(Self { + handle, + io, + phase: Phase::Opening { deadline }, + }) + } + Err(error) => { + io.finish(Err(error.to_string())); + None + } + } + } + + pub(super) fn handshake_deadline(&self) -> Option { + match self.phase { + Phase::Opening { deadline } => Some(deadline), + Phase::Open | Phase::ReceivedClose => None, + } + } + + pub(super) fn complete_handshake( + &mut self, + result: Result<(), curl::Error>, + ) -> Result { + if self.handshake_deadline().is_none() { + return Ok(Step::Idle); + } + let handshake = self.handle.get_mut(); + let result = match handshake.error.take() { + Some(error) => Err(error), + None => result.map_err(|error| error.to_string()), + }; + let event = CurlWebSocketEvent::Handshake { + request: std::mem::take(&mut handshake.request), + response: std::mem::take(&mut handshake.response), + result: result.clone(), + }; + // No data events precede the handshake, so this queue has a free slot. + let delivered = self.io.events.try_send(event).is_ok(); + result?; + if !delivered { + return Ok(Step::Closed); + } + self.phase = Phase::Open; + Ok(Step::Idle) + } + + pub(super) fn advance(&mut self) -> Result { + if self.io.cancelled() { + return Ok(Step::Closed); + } + if let Some(deadline) = self.handshake_deadline() { + return if deadline <= Instant::now() { + Err("WebSocket handshake timed out".to_owned()) + } else { + Ok(Step::Idle) + }; + } + let mut progressed = false; + for _ in 0..IO_BUDGET { + if self.io.cancelled() { + return Ok(Step::Closed); + } + // Receive delivery must never hold up a pending write. + progressed |= self.write_pending()?; + match self.read_chunk()? { + Step::Progress => progressed = true, + Step::Idle => break, + Step::Closed => return Ok(Step::Closed), + } + } + Ok(if progressed { + Step::Progress + } else { + Step::Idle + }) + } + + 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 { + return Ok(false); + }; + match self + .handle + .ws_send(&send.frame.data[send.offset..], 0, send.frame.flags) + { + Ok(count) => { + send.offset += count; + if send.offset == send.frame.data.len() { + let send = pending.take().expect("completed frame"); + let _ = send.completed.send(send.offset); + } + Ok(true) + } + Err(error) if error.is_again() => { + #[cfg(test)] + self.io.control.write_blocked.notify_one(); + Ok(false) + } + Err(error) => Err(format!("WebSocket send failed: {error}")), + } + } + + fn reading_enabled(&self) -> bool { + matches!(self.phase, Phase::Open) && self.io.control.reading.load(Ordering::Acquire) + } + + fn read_chunk(&mut self) -> Result { + if !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(_)) => { + #[cfg(test)] + self.io.control.read_blocked.notify_one(); + return Ok(Step::Idle); + } + Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => return Ok(Step::Closed), + }; + let mut data = vec![0; CHUNK_BYTES]; + match self.handle.ws_recv(&mut data) { + Ok((count, frame)) => { + data.truncate(count); + // Let the caller answer Close before probing EOF. A peer can + // half-close TCP in the same packet as its Close frame. + if frame.flags().contains(WsFlags::CLOSE) && frame.bytes_left() == 0 { + self.phase = Phase::ReceivedClose; + } + permit.send(CurlWebSocketEvent::Chunk { data, frame }); + Ok(Step::Progress) + } + Err(error) if error.is_again() => Ok(Step::Idle), + Err(error) if error.is_got_nothing() => Ok(Step::Closed), + Err(error) => Err(format!("WebSocket receive failed: {error}")), + } + } + + pub(super) fn wait_fd(&self) -> Option { + 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(); + if !reading && !writing { + return None; + } + let mut fd = WaitFd::new(); + fd.set_fd(self.handle.active_socket().ok()??); + fd.poll_on_read(reading).poll_on_write(writing); + Some(fd) + } + + 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); + self.io.finish(result); + } +}