use crate::courierust_body::Body;
use crate::courierust_bytes::Bytes;
use crate::courierust_error::{Error, ErrorKind, Result};
use crate::courierust_http::header::{HeaderMap, HeaderName, HeaderValue};
use crate::courierust_http::method::Method;
use crate::courierust_http::request::Request;
use crate::courierust_http::response::Response;
use crate::courierust_http::status::StatusCode;
use crate::courierust_http::version::Version;
use crate::courierust_io::BufReader;
use crate::courierust_net::ConnStream;
use crate::courierust_server::Handler;
use crate::courierust_ws::frame::{FrameSink, Mask, SharedSink};
use crate::courierust_ws::handshake::{
self, accept_key, CompressionParams, IpNet, OriginPolicy, PerMessageDeflate, PmDeflatePolicy,
WsOffer,
};
use crate::courierust_ws::session::{MaskSource, Role, Session, SessionConfig, Stats};
use crate::courierust_ws::writer::{CloseFlag, FrameWriter};
use crate::courierust_ws::Event;
use core::any::Any;
use core::net::IpAddr;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex, MutexGuard};
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct WsConfig {
pub enabled: bool,
pub origin: OriginPolicy,
pub trusted_proxies: Vec<IpNet>,
pub subprotocols: Vec<String>,
pub compression: PmDeflatePolicy,
pub max_frame: usize,
pub max_message: usize,
pub max_fragments: u32,
pub max_send_queue: usize,
pub read_buffer: usize,
pub ping_interval: Option<Duration>,
pub close_timeout: Option<Duration>,
}
impl Default for WsConfig {
fn default() -> Self {
Self {
enabled: true,
origin: OriginPolicy::default(),
trusted_proxies: Vec::new(),
subprotocols: Vec::new(),
compression: PmDeflatePolicy::default(),
max_frame: 16 * 1024 * 1024,
max_message: 16 * 1024 * 1024,
max_fragments: 0,
max_send_queue: 4 * 1024 * 1024,
read_buffer: 64 * 1024,
ping_interval: Some(Duration::from_secs(30)),
close_timeout: Some(Duration::from_secs(5)),
}
}
}
pub enum WsUpgradeReply {
Accept(Arc<dyn WsService>),
Refuse(Response<Body>),
Pass,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum WsData {
Text(String),
Binary(Bytes),
}
pub trait WsService: Send + Sync + 'static {
fn on_open(&self, _c: &mut WsConn) {}
fn on_message(&self, _c: &mut WsConn, _msg: WsData) {}
fn on_pong(&self, _c: &mut WsConn, _payload: &[u8]) {}
fn on_close(&self, _c: &mut WsConn, _code: Option<u16>, _clean: bool) {}
fn on_idle(&self, _c: &mut WsConn) {}
}
impl<F> WsService for F
where
F: Fn(&mut WsConn, WsData) + Send + Sync + 'static,
{
fn on_message(&self, c: &mut WsConn, msg: WsData) {
self(c, msg)
}
}
#[derive(Debug, Clone)]
pub struct WsInfo {
pub path: String,
pub peer: IpAddr,
pub client_ip: IpAddr,
pub secure: bool,
pub origin: Option<String>,
pub protocol: Option<String>,
pub compression: Option<PerMessageDeflate>,
}
pub struct OutQueue {
buf: Vec<u8>,
pos: usize,
limit: usize,
closed: bool,
}
impl OutQueue {
pub fn new(limit: usize) -> Self {
Self {
buf: Vec::new(),
pos: 0,
limit,
closed: false,
}
}
#[inline]
pub fn len(&self) -> usize {
self.buf.len() - self.pos
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len() == 0
}
#[inline]
pub fn is_closed(&self) -> bool {
self.closed
}
pub fn close(&mut self) {
self.closed = true;
self.buf.clear();
self.pos = 0;
}
pub fn append(&mut self, bytes: &[u8]) -> Result<()> {
if self.closed {
return Err(Error::canceled("websocket: connection is closed"));
}
self.reserve(bytes.len())?;
self.buf.extend_from_slice(bytes);
Ok(())
}
pub fn append_masked(&mut self, payload: &[u8], mask: Mask) -> Result<()> {
if self.closed {
return Err(Error::canceled("websocket: connection is closed"));
}
self.reserve(payload.len())?;
let start = self.buf.len();
self.buf.extend_from_slice(payload);
mask.apply(0, &mut self.buf[start..]);
Ok(())
}
fn reserve(&mut self, incoming: usize) -> Result<()> {
if self.pos > 0 && (self.pos == self.buf.len() || self.pos + incoming > self.limit) {
self.buf.drain(..self.pos);
self.pos = 0;
}
if self.buf.len() + incoming > self.limit {
self.closed = true;
self.buf.clear();
self.pos = 0;
return Err(Error::overflow("websocket: send queue overflow"));
}
Ok(())
}
pub fn drain(&mut self, writer: &mut impl crate::courierust_io::Write) -> Result<bool> {
while self.pos < self.buf.len() {
match crate::courierust_io::Write::write(writer, &self.buf[self.pos..]) {
Ok(0) => return Err(Error::io("websocket: write made no progress")),
Ok(n) => self.pos += n,
Err(e) if e.kind == ErrorKind::WouldBlock => return Ok(false),
Err(e) => return Err(e),
}
}
self.buf.clear();
self.pos = 0;
Ok(true)
}
}
#[derive(Clone)]
pub struct QueueSink {
queue: Arc<Mutex<OutQueue>>,
wake: Option<Arc<dyn Fn() + Send + Sync>>,
}
impl QueueSink {
pub fn new(queue: Arc<Mutex<OutQueue>>, wake: Option<Arc<dyn Fn() + Send + Sync>>) -> Self {
Self { queue, wake }
}
pub fn queue(&self) -> &Arc<Mutex<OutQueue>> {
&self.queue
}
}
impl FrameSink for QueueSink {
fn write_frame(&mut self, header: &[u8], payload: &[u8], mask: Option<Mask>) -> Result<()> {
{
let mut q = lock(&self.queue);
if q.len() + header.len() + payload.len() > q.limit {
q.close();
return Err(Error::overflow("websocket: send queue overflow"));
}
q.append(header)?;
match mask {
None => q.append(payload)?,
Some(m) => q.append_masked(payload, m)?,
}
}
if let Some(wake) = &self.wake {
wake();
}
Ok(())
}
}
pub(crate) type SharedStreamSink = SharedSink<Arc<ConnStream>>;
fn lock<T>(m: &Mutex<T>) -> MutexGuard<'_, T> {
m.lock().unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub type BoxSink = Box<dyn FrameSink + Send>;
struct WsConnInner {
info: WsInfo,
writer: Mutex<FrameWriter<BoxSink>>,
state: Mutex<Option<Box<dyn Any + Send>>>,
closing: AtomicBool,
}
#[derive(Clone)]
pub struct WsConn {
inner: Arc<WsConnInner>,
}
impl WsConn {
fn new(info: WsInfo, writer: FrameWriter<BoxSink>) -> Self {
Self {
inner: Arc::new(WsConnInner {
info,
writer: Mutex::new(writer),
state: Mutex::new(None),
closing: AtomicBool::new(false),
}),
}
}
pub fn info(&self) -> &WsInfo {
&self.inner.info
}
pub fn path(&self) -> &str {
&self.inner.info.path
}
pub fn client_ip(&self) -> IpAddr {
self.inner.info.client_ip
}
pub fn peer(&self) -> IpAddr {
self.inner.info.peer
}
pub fn is_secure(&self) -> bool {
self.inner.info.secure
}
pub fn protocol(&self) -> Option<&str> {
self.inner.info.protocol.as_deref()
}
pub fn compression(&self) -> Option<PerMessageDeflate> {
self.inner.info.compression
}
pub fn is_closing(&self) -> bool {
self.inner.closing.load(Ordering::Acquire)
}
pub fn set_state<T: Any + Send>(&self, value: T) {
let mut slot = lock(&self.inner.state);
*slot = Some(Box::new(value));
}
pub fn with_state<T: Any + Send, R>(&self, f: impl FnOnce(&mut T) -> R) -> Option<R> {
let mut slot = lock(&self.inner.state);
slot.as_mut()?.downcast_mut::<T>().map(f)
}
pub fn clear_state(&self) {
*lock(&self.inner.state) = None;
}
pub fn send_text(&self, text: &str) -> Result<()> {
lock(&self.inner.writer).send_text(text)
}
pub fn send_binary(&self, data: &[u8]) -> Result<()> {
lock(&self.inner.writer).send_binary(data)
}
pub fn send_ping(&self, payload: &[u8]) -> Result<()> {
lock(&self.inner.writer).send_ping(payload)
}
pub fn send_pong(&self, payload: &[u8]) -> Result<()> {
lock(&self.inner.writer).send_pong(payload)
}
pub fn close(&self, code: u16, reason: &str) -> Result<()> {
if self.inner.closing.swap(true, Ordering::AcqRel) {
return Ok(());
}
lock(&self.inner.writer).send_close(code, reason)
}
pub fn sender(&self) -> WsSender {
WsSender {
inner: self.inner.clone(),
}
}
pub fn stats(&self) -> Stats {
*lock(&self.inner.writer).stats()
}
}
#[derive(Clone)]
pub struct WsSender {
inner: Arc<WsConnInner>,
}
impl WsSender {
pub fn send_text(&self, text: &str) -> Result<()> {
lock(&self.inner.writer).send_text(text)
}
pub fn send_binary(&self, data: &[u8]) -> Result<()> {
lock(&self.inner.writer).send_binary(data)
}
pub fn send_ping(&self, payload: &[u8]) -> Result<()> {
lock(&self.inner.writer).send_ping(payload)
}
pub fn close(&self, code: u16, reason: &str) -> Result<()> {
if self.inner.closing.swap(true, Ordering::AcqRel) {
return Ok(());
}
lock(&self.inner.writer).send_close(code, reason)
}
pub fn is_closing(&self) -> bool {
self.inner.closing.load(Ordering::Acquire)
}
}
#[derive(Debug, Clone)]
pub struct WsRefusal {
pub status: StatusCode,
pub reason: &'static str,
pub advertise_version: bool,
}
impl WsRefusal {
fn new(status: u16, reason: &'static str) -> Self {
Self {
status: StatusCode::from_u16(status),
reason,
advertise_version: status == 426,
}
}
pub fn response(&self) -> Response<Body> {
let mut resp: Response<Body> = Response::with_status(self.status);
if self.advertise_version {
resp.headers.insert(
HeaderName::from_lowercase("sec-websocket-version"),
HeaderValue::from_static("13"),
);
}
resp.headers.insert(
HeaderName::from_lowercase("content-type"),
HeaderValue::from_static("text/plain; charset=utf-8"),
);
resp.headers.insert(
HeaderName::from_lowercase("connection"),
HeaderValue::from_static("close"),
);
resp.body = Body::Bytes(Bytes::from(alloc::format!("{}\n", self.reason)));
resp
}
}
#[derive(Debug, Clone)]
pub struct WsPlan {
pub offer: WsOffer,
pub protocol: Option<String>,
pub compression: Option<PerMessageDeflate>,
}
impl WsPlan {
pub fn server_compression(&self) -> Option<CompressionParams> {
self.compression.map(|p| p.server_view())
}
pub fn accept_headers(&self) -> Result<HeaderMap> {
let mut headers = HeaderMap::with_capacity(4);
headers.insert(
HeaderName::from_lowercase("upgrade"),
HeaderValue::from_static("websocket"),
);
headers.insert(
HeaderName::from_lowercase("connection"),
HeaderValue::from_static("Upgrade"),
);
headers.insert(
HeaderName::from_lowercase("sec-websocket-accept"),
HeaderValue::from_bytes(accept_key(&self.offer.key)?.as_bytes())?,
);
if let Some(protocol) = &self.protocol {
headers.insert(
HeaderName::from_lowercase("sec-websocket-protocol"),
HeaderValue::from_bytes(protocol.as_bytes())?,
);
}
if let Some(pm) = &self.compression {
headers.insert(
HeaderName::from_lowercase("sec-websocket-extensions"),
HeaderValue::from_bytes(pm.response_header().as_bytes())?,
);
}
Ok(headers)
}
}
pub fn plan(
req: &Request<Body>,
peer: IpAddr,
tls_active: bool,
ws: &WsConfig,
) -> core::result::Result<WsPlan, WsRefusal> {
if !ws.enabled {
return Err(WsRefusal::new(400, "websocket: disabled"));
}
if !handshake::is_websocket_upgrade(&req.headers) {
return Err(WsRefusal::new(400, "websocket: not an upgrade request"));
}
if req.method != Method::GET {
return Err(WsRefusal::new(405, "websocket: upgrade requires GET"));
}
if req.version != Version::HTTP_11 {
return Err(WsRefusal::new(400, "websocket: upgrade requires HTTP/1.1"));
}
let versions: Vec<&HeaderValue> = req.headers.get_all("sec-websocket-version").collect();
if versions.len() != 1 {
return Err(WsRefusal::new(
426,
"websocket: missing or duplicate Sec-WebSocket-Version",
));
}
let version_text = versions[0].to_str().unwrap_or("").trim();
if version_text != "13" {
return Err(WsRefusal::new(
426,
"websocket: only version 13 is supported",
));
}
let offer = match WsOffer::parse(req, peer, tls_active, &ws.trusted_proxies) {
Ok(o) => o,
Err(_) => return Err(WsRefusal::new(400, "websocket: malformed upgrade request")),
};
if !ws
.origin
.check(offer.origin.as_deref(), offer.request_origin().as_deref())
{
return Err(WsRefusal::new(403, "websocket: origin rejected"));
}
let protocol = if ws.subprotocols.is_empty() {
None
} else {
offer.select_protocol(&ws.subprotocols)
};
let compression = offer
.extension(crate::courierust_ws::handshake::PERMESSAGE_DEFLATE)
.and_then(|e| PerMessageDeflate::negotiate(e, &ws.compression));
Ok(WsPlan {
offer,
protocol,
compression,
})
}
pub fn session_config(ws: &WsConfig, params: Option<CompressionParams>) -> SessionConfig {
SessionConfig {
role: Role::Server,
max_frame: ws.max_frame,
max_message: ws.max_message,
max_fragments: ws.max_fragments,
compression: params,
auto_pong: true,
}
}
pub fn protocol_close(e: &Error) -> (u16, &'static str) {
match e.kind {
ErrorKind::Overflow => (1009, "message too big"),
ErrorKind::Protocol => {
if e.message
.as_deref()
.map(|m| m.contains("UTF-8"))
.unwrap_or(false)
{
(1007, "invalid payload data")
} else {
(1002, "protocol error")
}
}
_ => (1002, "protocol error"),
}
}
pub fn reports_violation(e: &Error) -> bool {
matches!(e.kind, ErrorKind::Protocol | ErrorKind::Overflow)
}
pub fn ask_handler(handler: &dyn Handler, req: &Request<Body>) -> WsUpgradeReply {
handler.websocket(req)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum WsStep {
Idle,
NeedWrite,
Close,
}
#[derive(Default)]
pub(crate) struct WakeSlot {
inner: Mutex<Option<Arc<dyn Fn() + Send + Sync>>>,
}
impl WakeSlot {
pub(crate) fn new() -> Arc<Self> {
Arc::new(Self::default())
}
pub(crate) fn set(&self, wake: Arc<dyn Fn() + Send + Sync>) {
*lock(&self.inner) = Some(wake);
}
pub(crate) fn fire(&self) {
let wake = lock(&self.inner).clone();
if let Some(wake) = wake {
wake();
}
}
}
pub(crate) struct WsEventConn {
socket: Arc<std::net::TcpStream>,
session: Session<Arc<std::net::TcpStream>, BoxSink>,
conn: WsConn,
queue: Arc<Mutex<OutQueue>>,
service: Arc<dyn WsService>,
ping_interval: Option<Duration>,
close_timeout: Option<Duration>,
last_recv: Instant,
last_ping: Instant,
closing_at: Option<Instant>,
reported: bool,
wake: Arc<WakeSlot>,
}
impl WsEventConn {
pub(crate) fn new(
socket: Arc<std::net::TcpStream>,
leftover: &[u8],
plan: WsPlan,
service: Arc<dyn WsService>,
ws: &WsConfig,
wake: Arc<WakeSlot>,
) -> Self {
let info = WsInfo {
path: plan.offer.path.clone(),
peer: plan.offer.peer,
client_ip: plan.offer.client_ip,
secure: plan.offer.secure,
origin: plan.offer.origin.clone(),
protocol: plan.protocol.clone(),
compression: plan.compression,
};
let params = plan.server_compression();
let queue = Arc::new(Mutex::new(OutQueue::new(ws.max_send_queue)));
let slot = wake.clone();
let sink = QueueSink::new(
queue.clone(),
Some(Arc::new(move || slot.fire()) as Arc<dyn Fn() + Send + Sync>),
);
let writer = FrameWriter::new(Box::new(sink.clone()) as BoxSink, MaskSource::None, params);
let cap = ws.read_buffer.max(16 * 1024) + leftover.len();
let mut reader = BufReader::new(socket.clone(), cap);
if !leftover.is_empty() {
reader.seed(leftover);
}
let session = Session::new(reader, writer, session_config(ws, params));
let conn_writer = FrameWriter::with_close_flag(
Box::new(sink) as BoxSink,
MaskSource::None,
params,
session.writer().close_flag(),
);
let mut conn = WsConn::new(info, conn_writer);
service.on_open(&mut conn);
Self {
socket,
session,
conn,
queue,
service,
ping_interval: ws.ping_interval,
close_timeout: ws.close_timeout,
last_recv: Instant::now(),
last_ping: Instant::now(),
closing_at: None,
reported: false,
wake,
}
}
pub(crate) fn set_wake(&mut self, wake: Arc<dyn Fn() + Send + Sync>) {
self.wake.set(wake);
}
pub(crate) fn socket(&self) -> &Arc<std::net::TcpStream> {
&self.socket
}
pub(crate) fn has_queued_output(&self) -> bool {
!lock(&self.queue).is_empty()
}
pub(crate) fn next_deadline(&self) -> Option<Instant> {
if let Some(started) = self.closing_at {
return Some(started + self.close_timeout.unwrap_or(Duration::from_secs(5)));
}
let interval = self.ping_interval?;
let close_due = self.last_recv + interval.saturating_mul(2);
let ping_due = core::cmp::max(self.last_recv + interval, self.last_ping + interval);
Some(core::cmp::min(ping_due, close_due))
}
fn try_flush(&mut self) -> core::result::Result<bool, ()> {
if !self.has_queued_output() {
return Ok(true);
}
let stream = self.socket.clone();
let mut writer: &std::net::TcpStream = &stream;
let drained = {
let mut q = lock(&self.queue);
q.drain(&mut writer)
};
match drained {
Ok(drained) => Ok(drained),
Err(_) => Err(()),
}
}
fn report_close(&mut self, code: Option<u16>, clean: bool) {
if !self.reported {
self.reported = true;
self.service.on_close(&mut self.conn, code, clean);
}
}
fn request_close(&mut self, code: u16, reason: &str) {
let _ = self.conn.close(code, reason);
if self.closing_at.is_none() {
self.closing_at = Some(Instant::now());
self.session.note_close_sent();
}
}
pub(crate) fn step(&mut self) -> WsStep {
match self.try_flush() {
Ok(true) => {}
Ok(false) => {
if lock(&self.queue).is_closed() {
self.request_close(1009, "send queue overflow");
return WsStep::Close;
}
return WsStep::NeedWrite;
}
Err(()) => {
self.report_close(None, false);
return WsStep::Close;
}
}
loop {
match self.session.poll_message() {
Ok(Some(event)) => {
self.last_recv = Instant::now();
match event {
Event::Text(t) => self.service.on_message(&mut self.conn, WsData::Text(t)),
Event::Binary(b) => {
self.service.on_message(&mut self.conn, WsData::Binary(b))
}
Event::Ping(_) => {
}
Event::Pong(p) => self.service.on_pong(&mut self.conn, &p),
Event::Close(frame) => {
let code = frame.as_ref().map(|f| f.code);
self.closing_at = Some(Instant::now());
self.report_close(code, true);
break;
}
}
if self.conn.is_closing() && self.closing_at.is_none() {
self.closing_at = Some(Instant::now());
self.session.note_close_sent();
}
}
Ok(None) => break,
Err(e) => {
if reports_violation(&e) {
let (code, reason) = protocol_close(&e);
self.request_close(code, reason);
self.report_close(Some(code), false);
break;
}
self.report_close(None, false);
return WsStep::Close;
}
}
}
if let Some(started) = self.closing_at {
match self.try_flush() {
Ok(false) => return WsStep::NeedWrite,
Err(()) => return WsStep::Close,
Ok(true) => {}
}
let timeout = self.close_timeout.unwrap_or(Duration::from_secs(5));
if self.session.close_received() || started.elapsed() >= timeout {
self.report_close(None, self.session.close_received());
return WsStep::Close;
}
return WsStep::Idle;
}
if let Some(interval) = self.ping_interval {
let idle = self.last_recv.elapsed();
if idle >= interval.saturating_mul(2) {
self.request_close(1001, "keepalive timeout");
self.report_close(Some(1001), false);
return WsStep::NeedWrite;
}
if idle >= interval {
let _ = self.conn.send_ping(b"");
self.last_ping = Instant::now();
self.service.on_idle(&mut self.conn);
}
}
if self.has_queued_output() {
return WsStep::NeedWrite;
}
WsStep::Idle
}
}
pub(crate) fn serve_blocking(
stream: Arc<ConnStream>,
reader: BufReader<Arc<ConnStream>>,
plan: WsPlan,
service: Arc<dyn WsService>,
ws: &WsConfig,
) -> Result<()> {
let info = WsInfo {
path: plan.offer.path.clone(),
peer: plan.offer.peer,
client_ip: plan.offer.client_ip,
secure: plan.offer.secure,
origin: plan.offer.origin.clone(),
protocol: plan.protocol.clone(),
compression: plan.compression,
};
let params = plan.server_compression();
let sink = SharedStreamSink::new(stream.clone());
let close_flag = CloseFlag::new();
let writer = FrameWriter::with_close_flag(
Box::new(sink.clone()) as BoxSink,
MaskSource::None,
params,
close_flag.clone(),
);
let mut session = Session::new(reader, writer, session_config(ws, params));
let conn_writer = FrameWriter::with_close_flag(
Box::new(sink) as BoxSink,
MaskSource::None,
params,
close_flag,
);
let mut conn = WsConn::new(info, conn_writer);
service.on_open(&mut conn);
let mut last_recv = Instant::now();
let mut close_code: Option<u16> = None;
let mut clean = false;
loop {
if let Some(interval) = ws.ping_interval {
let _ = stream.configure(Some(interval));
}
let polled = session.poll_message();
if ws.ping_interval.is_some() {
let _ = stream.configure(None);
}
match polled {
Ok(Some(event)) => {
last_recv = Instant::now();
match event {
Event::Text(t) => service.on_message(&mut conn, WsData::Text(t)),
Event::Binary(b) => service.on_message(&mut conn, WsData::Binary(b)),
Event::Ping(_) => {
}
Event::Pong(p) => service.on_pong(&mut conn, &p),
Event::Close(frame) => {
close_code = frame.as_ref().map(|f| f.code);
clean = true;
break;
}
}
if conn.is_closing() {
if let Some((code, is_clean)) = drain_close(&mut session, &stream, ws) {
close_code = code;
clean = is_clean;
}
break;
}
}
Ok(None) => {
break;
}
Err(e) => match e.kind {
ErrorKind::Timeout => {
let Some(interval) = ws.ping_interval else {
let _ = conn.close(1001, "idle timeout");
drain_close(&mut session, &stream, ws);
break;
};
if last_recv.elapsed() >= interval.saturating_mul(2) {
let _ = conn.close(1001, "keepalive timeout");
break;
}
if last_recv.elapsed() >= interval {
let _ = conn.send_ping(b"");
service.on_idle(&mut conn);
}
}
ErrorKind::UnexpectedEof => break,
_ => {
if reports_violation(&e) {
let (code, reason) = protocol_close(&e);
if !session.close_sent() {
let _ = session.close(code, reason);
}
close_code = Some(code);
}
break;
}
},
}
}
let _ = session.flush();
service.on_close(&mut conn, close_code, clean);
Ok(())
}
fn drain_close(
session: &mut Session<Arc<ConnStream>, BoxSink>,
stream: &Arc<ConnStream>,
ws: &WsConfig,
) -> Option<(Option<u16>, bool)> {
session.note_close_sent();
let timeout = ws.close_timeout?;
let _ = stream.configure(Some(timeout));
let deadline = Instant::now() + timeout;
loop {
if Instant::now() >= deadline {
return Some((None, false));
}
match session.poll_message() {
Ok(Some(Event::Close(frame))) => return Some((frame.as_ref().map(|f| f.code), true)),
Ok(Some(_)) => continue,
Ok(None) => return Some((None, false)),
Err(e) => match e.kind {
ErrorKind::Timeout => return Some((None, false)),
_ => return Some((None, false)),
},
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::courierust_http::header::HeaderName;
use crate::courierust_http::uri::PathAndQuery;
use crate::courierust_http::version::Version;
fn upgrade_request() -> Request<Body> {
let mut req = Request::new(Method::GET, PathAndQuery::from_static("/chat"));
req.version = Version::HTTP_11;
req.headers.append(
HeaderName::from_static("host"),
HeaderValue::from_static("ws.example.com"),
);
req.headers.append(
HeaderName::from_static("upgrade"),
HeaderValue::from_static("websocket"),
);
req.headers.append(
HeaderName::from_static("connection"),
HeaderValue::from_static("Upgrade"),
);
req.headers.append(
HeaderName::from_static("sec-websocket-key"),
HeaderValue::from_static("dGhlIHNhbXBsZSBub25jZQ=="),
);
req.headers.append(
HeaderName::from_static("sec-websocket-version"),
HeaderValue::from_static("13"),
);
req
}
#[test]
fn plan_accepts_a_well_formed_upgrade() {
let req = upgrade_request();
let p = plan(
&req,
"127.0.0.1".parse().unwrap(),
false,
&WsConfig::default(),
)
.unwrap();
assert_eq!(p.offer.path, "/chat");
let headers = p.accept_headers().unwrap();
assert_eq!(
headers
.get("sec-websocket-accept")
.unwrap()
.to_str()
.unwrap(),
"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="
);
assert_eq!(
headers.get("upgrade").unwrap().to_str().unwrap(),
"websocket"
);
assert_eq!(
headers.get("connection").unwrap().to_str().unwrap(),
"Upgrade"
);
}
#[test]
fn plan_refuses_wrong_version_with_426() {
let mut req = upgrade_request();
req.headers.remove("sec-websocket-version");
req.headers.append(
HeaderName::from_static("sec-websocket-version"),
HeaderValue::from_static("8"),
);
let refusal = plan(
&req,
"127.0.0.1".parse().unwrap(),
false,
&WsConfig::default(),
)
.unwrap_err();
assert_eq!(refusal.status, StatusCode::from_u16(426));
assert!(refusal.advertise_version);
let resp = refusal.response();
assert_eq!(
resp.headers
.get("sec-websocket-version")
.unwrap()
.to_str()
.unwrap(),
"13"
);
}
#[test]
fn plan_refuses_cross_origin_by_default() {
let mut req = upgrade_request();
req.headers.append(
HeaderName::from_static("origin"),
HeaderValue::from_static("https://evil.test"),
);
let refusal = plan(
&req,
"127.0.0.1".parse().unwrap(),
false,
&WsConfig::default(),
)
.unwrap_err();
assert_eq!(refusal.status, StatusCode::from_u16(403));
let mut req = upgrade_request();
req.headers.append(
HeaderName::from_static("origin"),
HeaderValue::from_static("http://ws.example.com"),
);
assert!(plan(
&req,
"127.0.0.1".parse().unwrap(),
false,
&WsConfig::default()
)
.is_ok());
}
#[test]
fn plan_refuses_a_duplicated_version_header() {
let mut req = upgrade_request();
req.headers.append(
HeaderName::from_static("sec-websocket-version"),
HeaderValue::from_static("13"),
);
assert!(plan(
&req,
"127.0.0.1".parse().unwrap(),
false,
&WsConfig::default()
)
.is_err());
}
#[test]
fn plan_is_inert_when_websockets_are_disabled() {
let req = upgrade_request();
let ws = WsConfig {
enabled: false,
..Default::default()
};
let refusal = plan(&req, "127.0.0.1".parse().unwrap(), false, &ws).unwrap_err();
assert_eq!(refusal.status, StatusCode::from_u16(400));
}
#[test]
fn the_application_writer_stops_after_a_close_on_the_session() {
let info = WsInfo {
path: String::from("/x"),
peer: "127.0.0.1".parse().unwrap(),
client_ip: "127.0.0.1".parse().unwrap(),
secure: false,
origin: None,
protocol: None,
compression: None,
};
let flag = CloseFlag::new();
let queue = Arc::new(Mutex::new(OutQueue::new(4096)));
let sink = QueueSink::new(queue.clone(), None);
let app_writer = FrameWriter::with_close_flag(
Box::new(sink.clone()) as BoxSink,
MaskSource::None,
None,
flag.clone(),
);
let conn = WsConn::new(info, app_writer);
conn.send_text("hello").unwrap();
let mut session_writer =
FrameWriter::with_close_flag(Box::new(sink) as BoxSink, MaskSource::None, None, flag);
session_writer
.send_close(crate::courierust_ws::close::NORMAL, "bye")
.unwrap();
assert!(
!conn.is_closing(),
"the close was the writer's, not the application handle's"
);
let err = conn.send_text("too late").unwrap_err();
assert_eq!(err.kind, ErrorKind::Canceled, "{err}");
assert!(conn.sender().send_binary(b"too late").is_err());
let queued = lock(&queue).buf.clone();
let mut frames = 0usize;
let mut pos = 0usize;
while pos < queued.len() {
let header = crate::courierust_ws::FrameHeader::parse(&queued[pos..])
.unwrap()
.unwrap();
pos += header.header_len + header.payload_len as usize;
frames += 1;
}
assert_eq!(frames, 2);
}
#[test]
fn subprotocol_selection_follows_server_preference() {
let mut req = upgrade_request();
req.headers.append(
HeaderName::from_static("sec-websocket-protocol"),
HeaderValue::from_static("chat.v1, chat.v2"),
);
let ws = WsConfig {
subprotocols: alloc::vec!["chat.v2".to_string(), "chat.v1".to_string()],
..Default::default()
};
let p = plan(&req, "127.0.0.1".parse().unwrap(), false, &ws).unwrap();
assert_eq!(p.protocol.as_deref(), Some("chat.v2"));
let headers = p.accept_headers().unwrap();
assert_eq!(
headers
.get("sec-websocket-protocol")
.unwrap()
.to_str()
.unwrap(),
"chat.v2"
);
}
#[test]
fn extension_negotiation_is_reflected_in_the_accept_headers() {
let mut req = upgrade_request();
req.headers.append(
HeaderName::from_static("sec-websocket-extensions"),
HeaderValue::from_static("permessage-deflate; client_max_window_bits=10"),
);
let p = plan(
&req,
"127.0.0.1".parse().unwrap(),
false,
&WsConfig::default(),
)
.unwrap();
let pm = p.compression.expect("permessage-deflate selected");
assert_eq!(pm.client_max_window_bits, 10);
let headers = p.accept_headers().unwrap();
assert!(headers
.get("sec-websocket-extensions")
.unwrap()
.to_str()
.unwrap()
.starts_with("permessage-deflate"));
let params = p.server_compression().unwrap();
assert_eq!(params.recv_window_bits, 10);
assert_eq!(params.send_window_bits, 15);
}
#[test]
fn queue_overflow_fails_the_frame_atomically() {
let q = Arc::new(Mutex::new(OutQueue::new(16)));
let mut sink = QueueSink::new(q.clone(), None);
assert!(sink.write_frame(&[0u8; 4], &[0u8; 8], None).is_ok());
assert!(sink.write_frame(&[0u8; 4], &[0u8; 8], None).is_err());
let guard = lock(&q);
assert!(guard.is_closed());
assert!(guard.is_empty());
}
#[test]
fn queue_drains_into_a_writer_and_compacts() {
let q = Arc::new(Mutex::new(OutQueue::new(64)));
let mut sink = QueueSink::new(q.clone(), None);
for _ in 0..4 {
sink.write_frame(&[0xa1, 0x01], &[0xff], None).unwrap();
}
let mut out = crate::courierust_io::VecWriter(Vec::new());
assert!(lock(&q).drain(&mut out).unwrap());
assert_eq!(out.0.len(), 12);
sink.write_frame(&[0xa1, 0x01], &[0xff], None).unwrap();
assert_eq!(lock(&q).len(), 3);
}
#[test]
fn protocol_errors_map_to_the_right_close_codes() {
assert_eq!(protocol_close(&Error::overflow("x")).0, 1009);
assert_eq!(
protocol_close(&Error::protocol(
"websocket: invalid UTF-8 in a text message"
))
.0,
1007
);
assert_eq!(protocol_close(&Error::protocol("anything else")).0, 1002);
}
#[test]
fn only_a_peer_violation_is_reported_to_the_peer() {
assert!(reports_violation(&Error::protocol("websocket: bad frame")));
assert!(reports_violation(&Error::overflow("websocket: too big")));
assert!(reports_violation(&Error::protocol(
"websocket: invalid UTF-8"
)));
assert!(!reports_violation(&Error::io("connection reset")));
assert!(!reports_violation(&Error::eof()));
assert!(!reports_violation(&Error::canceled("gone")));
assert!(!reports_violation(&Error::with_message(
ErrorKind::Timeout,
"no data"
)));
}
}