use std::fmt;
use std::io;
use crate::error::{Error, Result};
use crate::io::runtime::{AsyncConn, Runtime};
use crate::url::Url;
use crate::websocket::{
base64_encode, build_client_frame, close_payload, derive_accept, parse_close_payload,
random_16, try_parse_frame, validate_control_frame, Frame, OPCODE_BINARY, OPCODE_CLOSE,
OPCODE_CONT, OPCODE_PING, OPCODE_PONG, OPCODE_TEXT,
};
use super::WsMessage;
#[cfg(any(feature = "rustls-tls", feature = "purecrypto-tls"))]
use crate::io::asynctls::AsyncTlsStream;
const MAX_HANDSHAKE_HEAD: usize = 64 * 1024;
enum Transport<C> {
Plain(C),
#[cfg(any(feature = "rustls-tls", feature = "purecrypto-tls"))]
Tls(Box<AsyncTlsStream<C>>),
}
impl<C: AsyncConn> AsyncConn for Transport<C> {
async fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
match self {
Transport::Plain(c) => c.read(buf).await,
#[cfg(any(feature = "rustls-tls", feature = "purecrypto-tls"))]
Transport::Tls(t) => t.read(buf).await,
}
}
async fn write_all(&mut self, buf: &[u8]) -> io::Result<()> {
match self {
Transport::Plain(c) => c.write_all(buf).await,
#[cfg(any(feature = "rustls-tls", feature = "purecrypto-tls"))]
Transport::Tls(t) => t.write_all(buf).await,
}
}
async fn flush(&mut self) -> io::Result<()> {
match self {
Transport::Plain(c) => c.flush().await,
#[cfg(any(feature = "rustls-tls", feature = "purecrypto-tls"))]
Transport::Tls(t) => t.flush().await,
}
}
}
pub struct WebSocket<C> {
transport: Transport<C>,
rxbuf: Vec<u8>,
closed: bool,
subprotocol: Option<String>,
}
impl<C> fmt::Debug for WebSocket<C> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WebSocket")
.field("closed", &self.closed)
.field("subprotocol", &self.subprotocol)
.field("buffered", &self.rxbuf.len())
.finish_non_exhaustive()
}
}
impl<C: AsyncConn> WebSocket<C> {
pub async fn connect<R>(rt: &R, url: &str) -> Result<WebSocket<C>>
where
R: Runtime<Conn = C>,
{
Self::connect_with_subprotocols(rt, url, &[]).await
}
pub async fn connect_with_subprotocols<R>(
rt: &R,
url: &str,
subprotocols: &[&str],
) -> Result<WebSocket<C>>
where
R: Runtime<Conn = C>,
{
let u = Url::parse(url)?;
let conn = super::native::connect(rt, &u.host, u.port).await?;
let transport = match u.scheme.as_str() {
"ws" => Transport::Plain(conn),
"wss" => {
#[cfg(any(feature = "rustls-tls", feature = "purecrypto-tls"))]
{
let mut opts = crate::tls::TlsOpts::verifying();
let tls = AsyncTlsStream::connect(conn, &u.host, &mut opts).await?;
Transport::Tls(Box::new(tls))
}
#[cfg(not(any(feature = "rustls-tls", feature = "purecrypto-tls")))]
{
let _ = conn;
return Err(Error::UnsupportedScheme(
"wss (no TLS backend compiled)".into(),
));
}
}
other => return Err(Error::UnsupportedScheme(other.to_string())),
};
let subs: Vec<String> = subprotocols.iter().map(|s| s.to_string()).collect();
let mut ws = WebSocket {
transport,
rxbuf: Vec::new(),
closed: false,
subprotocol: None,
};
ws.handshake(&u, &subs).await?;
Ok(ws)
}
pub fn subprotocol(&self) -> Option<&str> {
self.subprotocol.as_deref()
}
pub fn is_closed(&self) -> bool {
self.closed
}
pub async fn recv(&mut self) -> Option<Result<WsMessage>> {
match self.recv_inner().await {
Ok(Some(m)) => Some(Ok(m)),
Ok(None) => None,
Err(e) => Some(Err(e)),
}
}
pub async fn send_text(&mut self, text: &str) -> Result<()> {
self.send_data(OPCODE_TEXT, text.as_bytes()).await
}
pub async fn send_binary(&mut self, data: &[u8]) -> Result<()> {
self.send_data(OPCODE_BINARY, data).await
}
pub async fn send(&mut self, msg: &WsMessage) -> Result<()> {
match msg {
WsMessage::Text(t) => self.send_text(t).await,
WsMessage::Binary(b) => self.send_binary(b).await,
}
}
pub async fn close(&mut self) -> Result<()> {
self.send_close(&[]).await
}
pub async fn close_with(&mut self, code: u16, reason: &str) -> Result<()> {
let payload = close_payload(code, reason)?;
self.send_close(&payload).await
}
async fn send_close(&mut self, payload: &[u8]) -> Result<()> {
if self.closed {
return Ok(());
}
self.closed = true;
let frame = build_client_frame(OPCODE_CLOSE, payload)?;
self.transport.write_all(&frame).await.map_err(Error::Io)?;
self.transport.flush().await.map_err(Error::Io)?;
Ok(())
}
async fn send_data(&mut self, opcode: u8, payload: &[u8]) -> Result<()> {
let frame = build_client_frame(opcode, payload)?;
self.transport.write_all(&frame).await.map_err(Error::Io)?;
self.transport.flush().await.map_err(Error::Io)?;
Ok(())
}
async fn recv_inner(&mut self) -> Result<Option<WsMessage>> {
if self.closed {
return Ok(None);
}
let mut assembled: Vec<u8> = Vec::new();
let mut msg_opcode: Option<u8> = None;
loop {
let frame = self.next_frame().await?;
match frame.opcode {
OPCODE_PING => {
validate_control_frame(&frame)?;
let pong = build_client_frame(OPCODE_PONG, &frame.payload)?;
self.transport.write_all(&pong).await.map_err(Error::Io)?;
self.transport.flush().await.map_err(Error::Io)?;
}
OPCODE_PONG => validate_control_frame(&frame)?,
OPCODE_CLOSE => {
validate_control_frame(&frame)?;
let _ = parse_close_payload(&frame.payload); if let Ok(close) = build_client_frame(OPCODE_CLOSE, &[]) {
let _ = self.transport.write_all(&close).await;
let _ = self.transport.flush().await;
}
self.closed = true;
return Ok(None);
}
OPCODE_TEXT | OPCODE_BINARY => {
if msg_opcode.is_some() {
return Err(Error::BadResponse(
"new data frame began before the previous message finished".into(),
));
}
if frame.rsv1 {
return Err(Error::BadResponse(
"RSV1 set but permessage-deflate was not negotiated".into(),
));
}
msg_opcode = Some(frame.opcode);
assembled.extend_from_slice(&frame.payload);
if frame.fin {
return Ok(Some(build_message(frame.opcode, assembled)?));
}
}
OPCODE_CONT => {
let op = msg_opcode.ok_or_else(|| {
Error::BadResponse("CONTINUATION frame without a start frame".into())
})?;
assembled.extend_from_slice(&frame.payload);
if frame.fin {
return Ok(Some(build_message(op, assembled)?));
}
}
other => return Err(Error::BadResponse(format!("unknown WS opcode 0x{other:x}"))),
}
}
}
async fn next_frame(&mut self) -> Result<Frame> {
let mut tmp = [0u8; 16 * 1024];
loop {
if let Some((frame, consumed)) = try_parse_frame(&self.rxbuf)? {
self.rxbuf.drain(..consumed);
return Ok(frame);
}
let n = self.transport.read(&mut tmp).await.map_err(Error::Io)?;
if n == 0 {
return Err(Error::UnexpectedEof);
}
self.rxbuf.extend_from_slice(&tmp[..n]);
}
}
async fn handshake(&mut self, u: &Url, subprotocols: &[String]) -> Result<()> {
let key_b64 = base64_encode(&random_16()?);
let host_header =
if (u.scheme == "ws" && u.port == 80) || (u.scheme == "wss" && u.port == 443) {
u.host.clone()
} else {
format!("{}:{}", u.host, u.port)
};
let path = if u.path.is_empty() {
"/"
} else {
u.path.as_str()
};
let proto_header = if subprotocols.is_empty() {
String::new()
} else {
format!("Sec-WebSocket-Protocol: {}\r\n", subprotocols.join(", "))
};
let req = format!(
"GET {path} HTTP/1.1\r\n\
Host: {host_header}\r\n\
Upgrade: websocket\r\n\
Connection: Upgrade\r\n\
Sec-WebSocket-Key: {key_b64}\r\n\
Sec-WebSocket-Version: 13\r\n\
{proto_header}\
\r\n"
);
self.transport
.write_all(req.as_bytes())
.await
.map_err(Error::Io)?;
self.transport.flush().await.map_err(Error::Io)?;
let head = self.read_handshake_head().await?;
self.interpret_handshake(&head, &key_b64)
}
async fn read_handshake_head(&mut self) -> Result<Vec<u8>> {
let mut tmp = [0u8; 4096];
loop {
if let Some(end) = find_double_crlf(&self.rxbuf) {
return Ok(self.rxbuf.drain(..end).collect());
}
let n = self.transport.read(&mut tmp).await.map_err(Error::Io)?;
if n == 0 {
return Err(Error::UnexpectedEof);
}
self.rxbuf.extend_from_slice(&tmp[..n]);
if self.rxbuf.len() > MAX_HANDSHAKE_HEAD {
return Err(Error::BadResponse(
"websocket handshake response headers too large".into(),
));
}
}
}
fn interpret_handshake(&mut self, head: &[u8], key_b64: &str) -> Result<()> {
let text = std::str::from_utf8(head)
.map_err(|_| Error::BadResponse("non-utf8 handshake response".into()))?;
let mut lines = text.split("\r\n");
let status = lines
.next()
.ok_or_else(|| Error::BadResponse("empty handshake response".into()))?;
if !(status.starts_with("HTTP/1.1 101") || status.starts_with("HTTP/1.0 101")) {
return Err(Error::BadResponse(format!(
"expected 101 Switching Protocols, got: {status:?}"
)));
}
let mut upgrade_ok = false;
let mut connection_ok = false;
let mut accept_value: Option<String> = None;
let mut subprotocol_value: Option<String> = None;
let mut had_extension = false;
for line in lines {
if line.is_empty() {
break;
}
let (k, v) = match line.split_once(':') {
Some((k, v)) => (k.trim(), v.trim()),
None => continue,
};
if k.eq_ignore_ascii_case("upgrade") {
if v.eq_ignore_ascii_case("websocket") {
upgrade_ok = true;
}
} else if k.eq_ignore_ascii_case("connection") {
if v.split(',')
.any(|t| t.trim().eq_ignore_ascii_case("upgrade"))
{
connection_ok = true;
}
} else if k.eq_ignore_ascii_case("sec-websocket-accept") {
accept_value = Some(v.to_string());
} else if k.eq_ignore_ascii_case("sec-websocket-protocol") {
subprotocol_value = Some(v.to_string());
} else if k.eq_ignore_ascii_case("sec-websocket-extensions") && !v.is_empty() {
had_extension = true;
}
}
if !upgrade_ok {
return Err(Error::BadResponse(
"missing or wrong Upgrade header in handshake response".into(),
));
}
if !connection_ok {
return Err(Error::BadResponse(
"missing or wrong Connection header in handshake response".into(),
));
}
let accept = accept_value
.ok_or_else(|| Error::BadResponse("missing Sec-WebSocket-Accept header".into()))?;
if accept != derive_accept(key_b64) {
return Err(Error::BadResponse("Sec-WebSocket-Accept mismatch".into()));
}
if had_extension {
return Err(Error::BadResponse(
"server negotiated a permessage extension that was not offered".into(),
));
}
self.subprotocol = subprotocol_value;
Ok(())
}
}
fn build_message(opcode: u8, payload: Vec<u8>) -> Result<WsMessage> {
if opcode == OPCODE_TEXT {
let s = String::from_utf8(payload)
.map_err(|_| Error::BadResponse("invalid UTF-8 in text message".into()))?;
Ok(WsMessage::Text(s))
} else {
Ok(WsMessage::Binary(payload))
}
}
fn find_double_crlf(buf: &[u8]) -> Option<usize> {
buf.windows(4).position(|w| w == b"\r\n\r\n").map(|i| i + 4)
}