refactor(websocket): encapsulate native session lifecycle

This commit is contained in:
ldm0
2026-09-11 01:47:10 +08:00
committed by Donough Liu
parent e3f8765a34
commit cd3a2b1679
3 changed files with 328 additions and 208 deletions
+1
View File
@@ -6,6 +6,7 @@
mod owner;
mod request;
mod session;
#[cfg(test)]
mod tests;
+102 -208
View File
@@ -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<Handshake>,
deadline: Instant,
open: bool,
received_close: bool,
io: SessionIo,
struct Owner {
multi: Multi,
sessions: HashMap<CurlTransferId, Session>,
dns: CurlDnsOwnerResidence<CurlTransferId, Pending>,
}
pub(super) fn run(
@@ -43,11 +43,30 @@ pub(super) fn run(
waker_tx: crossbeam_channel::Sender<MultiWaker>,
shutdown: Arc<AtomicBool>,
) {
let mut multi = Multi::new();
let _ = waker_tx.send(multi.waker());
let mut sessions = HashMap::<CurlTransferId, Session>::new();
let mut dns = CurlDnsOwnerResidence::<CurlTransferId, Pending>::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<Submission>) {
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<Step, String>) {
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<CurlTransferId, Session>, 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<bool, std::result::Result<(), String>> {
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<WaitFd> {
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)
}
}
+225
View File
@@ -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<Handshake>,
io: SessionIo,
phase: Phase,
}
impl Session {
pub(super) fn attach(
multi: &mut Multi,
id: CurlTransferId,
mut easy: Easy2<Handshake>,
io: SessionIo,
deadline: Instant,
) -> Option<Self> {
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<Instant> {
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<Step, String> {
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<Step, String> {
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<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 {
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<Step, String> {
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<WaitFd> {
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);
}
}