use std::io::{Read, Write};
use std::net::TcpStream;
use tungstenite::handshake::machine::TryParse;
use tungstenite::handshake::server::{Request, create_response, write_response};
use tungstenite::http::{HeaderValue, Response as HttpResponse, StatusCode};
use tungstenite::protocol::WebSocketConfig as TungsteniteConfig;
use liminal::protocol::{Frame, ProtocolError, decode};
use liminal_protocol::wire::FRAME_MAX;
use crate::ServerError;
#[path = "websocket/listener.rs"]
mod listener;
#[path = "websocket/outbound.rs"]
pub(super) mod outbound;
#[path = "websocket/process.rs"]
mod process;
#[path = "websocket/supervisor.rs"]
mod supervisor;
#[cfg(test)]
#[path = "websocket/handshake_tests.rs"]
mod handshake_tests;
pub use listener::WebSocketListener;
pub(super) const MAX_UPGRADE_REQUEST_BYTES: usize = 8192;
pub(super) fn liminal_ws_message_bound() -> Result<usize, ServerError> {
usize::try_from(FRAME_MAX).map_err(|_| ServerError::ListenerAccept {
message: format!(
"websocket acceptor cannot start: this target's usize cannot represent the \
liminal frame bound of {FRAME_MAX} bytes"
),
})
}
pub(super) fn pinned_protocol_config(message_bound: usize) -> TungsteniteConfig {
let base = TungsteniteConfig::default();
let write_headroom = base.write_buffer_size.saturating_mul(2);
base.max_message_size(Some(message_bound))
.max_frame_size(Some(message_bound))
.max_write_buffer_size(message_bound.saturating_add(write_headroom))
}
#[derive(Debug)]
pub(super) struct AcceptorSettings {
pub(super) path: String,
pub(super) allowed_origins: Vec<String>,
pub(super) ping_interval: Option<std::time::Duration>,
pub(super) message_bound: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(super) enum UpgradeRefusal {
OversizedRequestHead {
received: usize,
},
MalformedRequest {
detail: String,
},
JunkAfterRequest,
NotAWebSocketUpgrade {
detail: String,
},
WrongPath {
requested: String,
},
DuplicateOriginHeader,
MalformedOriginHeader,
OriginNotAllowed {
origin: String,
},
SubprotocolOffered,
}
impl UpgradeRefusal {
pub(super) const fn status(&self) -> StatusCode {
match self {
Self::OversizedRequestHead { .. } => StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE,
Self::MalformedRequest { .. }
| Self::JunkAfterRequest
| Self::NotAWebSocketUpgrade { .. }
| Self::DuplicateOriginHeader
| Self::MalformedOriginHeader
| Self::SubprotocolOffered => StatusCode::BAD_REQUEST,
Self::WrongPath { .. } => StatusCode::NOT_FOUND,
Self::OriginNotAllowed { .. } => StatusCode::FORBIDDEN,
}
}
}
impl std::fmt::Display for UpgradeRefusal {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::OversizedRequestHead { received } => write!(
formatter,
"upgrade request head exceeded the {MAX_UPGRADE_REQUEST_BYTES}-byte bound \
({received} bytes received)"
),
Self::MalformedRequest { detail } => {
write!(formatter, "malformed HTTP request: {detail}")
}
Self::JunkAfterRequest => {
write!(
formatter,
"bytes followed the request head before the upgrade"
)
}
Self::NotAWebSocketUpgrade { detail } => {
write!(formatter, "not a well-formed WebSocket upgrade: {detail}")
}
Self::WrongPath { requested } => {
write!(
formatter,
"request path '{requested}' is not the upgrade path"
)
}
Self::DuplicateOriginHeader => {
write!(formatter, "multiple Origin headers present")
}
Self::MalformedOriginHeader => {
write!(formatter, "Origin header is not valid visible ASCII")
}
Self::OriginNotAllowed { origin } => {
write!(
formatter,
"origin '{origin}' is not on the configured allow-list"
)
}
Self::SubprotocolOffered => {
write!(formatter, "subprotocols are not negotiated on this route")
}
}
}
}
#[derive(Debug)]
pub(super) enum HandshakeOutcome {
Upgraded,
Refused(UpgradeRefusal),
SocketError(std::io::Error),
}
pub(super) fn perform_upgrade(
stream: &mut TcpStream,
settings: &AcceptorSettings,
) -> HandshakeOutcome {
let request = match read_upgrade_request(stream) {
Ok(Ok(request)) => request,
Ok(Err(refusal)) => return refuse(stream, refusal),
Err(error) => return HandshakeOutcome::SocketError(error),
};
match validate_upgrade_request(&request, settings) {
Ok(()) => {}
Err(refusal) => return refuse(stream, refusal),
}
let response = match create_response(&request) {
Ok(response) => response,
Err(error) => {
return refuse(
stream,
UpgradeRefusal::NotAWebSocketUpgrade {
detail: error.to_string(),
},
);
}
};
let mut serialized: Vec<u8> = Vec::new();
if let Err(error) = write_response(&mut serialized, &response) {
return HandshakeOutcome::SocketError(std::io::Error::other(error.to_string()));
}
if let Err(error) = stream.write_all(&serialized).and_then(|()| stream.flush()) {
return HandshakeOutcome::SocketError(error);
}
HandshakeOutcome::Upgraded
}
fn read_upgrade_request(
stream: &mut TcpStream,
) -> Result<Result<Request, UpgradeRefusal>, std::io::Error> {
let mut head: Vec<u8> = Vec::new();
let mut chunk = [0_u8; 1024];
loop {
match Request::try_parse(&head) {
Ok(Some((parsed_length, request))) => {
if parsed_length != head.len() {
return Ok(Err(UpgradeRefusal::JunkAfterRequest));
}
return Ok(Ok(request));
}
Ok(None) => {}
Err(error) => {
return Ok(Err(UpgradeRefusal::MalformedRequest {
detail: error.to_string(),
}));
}
}
if head.len() >= MAX_UPGRADE_REQUEST_BYTES {
return Ok(Err(UpgradeRefusal::OversizedRequestHead {
received: head.len(),
}));
}
let read = match stream.read(&mut chunk) {
Ok(0) => {
return Err(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"peer closed before completing the upgrade request head",
));
}
Ok(read) => read,
Err(error) if error.kind() == std::io::ErrorKind::Interrupted => continue,
Err(error) => return Err(error),
};
head.extend_from_slice(chunk.get(..read).unwrap_or(&[]));
if head.len() > MAX_UPGRADE_REQUEST_BYTES {
return Ok(Err(UpgradeRefusal::OversizedRequestHead {
received: head.len(),
}));
}
}
}
pub(super) fn validate_upgrade_request(
request: &Request,
settings: &AcceptorSettings,
) -> Result<(), UpgradeRefusal> {
if request.uri().path() != settings.path || request.uri().query().is_some() {
return Err(UpgradeRefusal::WrongPath {
requested: request.uri().to_string(),
});
}
let mut origins = request.headers().get_all("Origin").iter();
let origin = origins.next();
if origins.next().is_some() {
return Err(UpgradeRefusal::DuplicateOriginHeader);
}
if let Some(origin) = origin {
let Ok(origin) = origin.to_str() else {
return Err(UpgradeRefusal::MalformedOriginHeader);
};
if !settings
.allowed_origins
.iter()
.any(|allowed| allowed == origin)
{
return Err(UpgradeRefusal::OriginNotAllowed {
origin: origin.to_owned(),
});
}
}
if request.headers().get("Sec-WebSocket-Protocol").is_some() {
return Err(UpgradeRefusal::SubprotocolOffered);
}
Ok(())
}
fn refuse(stream: &mut TcpStream, refusal: UpgradeRefusal) -> HandshakeOutcome {
let response = HttpResponse::builder()
.status(refusal.status())
.header("Connection", HeaderValue::from_static("close"))
.header("Content-Length", HeaderValue::from_static("0"))
.body(());
match response {
Ok(response) => {
let mut serialized: Vec<u8> = Vec::new();
if let Err(error) = write_response(&mut serialized, &response) {
tracing::debug!(%error, "websocket refusal response could not be serialized");
} else if let Err(error) = stream.write_all(&serialized).and_then(|()| stream.flush()) {
tracing::debug!(%error, "websocket refusal response could not be written");
}
}
Err(error) => {
tracing::debug!(%error, "websocket refusal response could not be constructed");
}
}
if let Err(error) = stream.shutdown(std::net::Shutdown::Both) {
tracing::debug!(%error, "websocket refusal socket shutdown failed");
}
HandshakeOutcome::Refused(refusal)
}
#[derive(Debug)]
pub(super) enum WsInboundViolation {
TextMessage,
MalformedFrame(ProtocolError),
TrailingBytes {
consumed: usize,
length: usize,
},
}
impl std::fmt::Display for WsInboundViolation {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::TextMessage => write!(
formatter,
"text message on the canonical binary-only websocket route"
),
Self::MalformedFrame(error) => write!(
formatter,
"binary message is not one complete canonical liminal frame: {error}"
),
Self::TrailingBytes { consumed, length } => write!(
formatter,
"binary message carried trailing bytes: frame consumed {consumed} of {length}"
),
}
}
}
pub(super) fn decode_ws_binary(bytes: &[u8]) -> Result<Frame, WsInboundViolation> {
let (frame, consumed) = decode(bytes).map_err(WsInboundViolation::MalformedFrame)?;
if consumed != bytes.len() {
return Err(WsInboundViolation::TrailingBytes {
consumed,
length: bytes.len(),
});
}
Ok(frame)
}