use std::io;
use crate::error::ErrorDiagnostic;
use crate::http::error::{DecodeError, EncodeError, ResponseError};
use crate::http::header::{ALLOW, HeaderValue, SEC_WEBSOCKET_VERSION};
use crate::http::{Response, StatusCode};
use crate::{connect::ConnectError, util::Either, util::clone_io_error};
use super::OpCode;
#[derive(Debug, thiserror::Error)]
pub enum WsError<E> {
#[error("Service error")]
Service(#[source] E),
#[error("Keep-alive error")]
KeepAlive,
#[error("Frame read timeout")]
ReadTimeout,
#[error("Write timeout")]
WriteTimeout,
#[error("Ws protocol level error")]
Protocol(#[source] ProtocolError),
#[error("Ws handshake error")]
Handshake(#[from] HandshakeError),
#[error("Peer has been disconnected: {0:?}")]
Disconnected(#[source] Option<io::Error>),
}
#[derive(Copy, Clone, Debug, thiserror::Error)]
pub enum ProtocolError {
#[error("Received an unmasked frame from client")]
UnmaskedFrame,
#[error("Received a masked frame from server")]
MaskedFrame,
#[error("Invalid opcode: {0}")]
InvalidOpcode(u8),
#[error("Reserved frame bits are set: {0:#05b}")]
ReservedBits(u8),
#[error("Fragmented control frame: {0}")]
FragmentedControlFrame(OpCode),
#[error("Invalid control frame length: {0}")]
InvalidLength(usize),
#[error("Invalid payload length encoding")]
InvalidLengthEncoding,
#[error("Invalid close status code: {0}")]
InvalidCloseCode(u16),
#[error("Invalid close-frame payload")]
InvalidClosePayload,
#[error("Invalid UTF-8 in close-frame description")]
InvalidUtf8,
#[error("WebSocket codec is closed")]
Closed,
#[error("A payload reached size limit.")]
Overflow,
#[error("Continuation is not started.")]
ContinuationNotStarted,
#[error("Received new continuation but it is already started")]
ContinuationStarted,
}
#[derive(Clone, Debug, thiserror::Error)]
pub enum WsConfigError {
#[error("Missing url scheme")]
MissingScheme,
#[error("Unknown url scheme")]
UnknownScheme,
#[error("Missing host name")]
MissingHost,
#[error("Url parse error: {0}")]
Parse(
#[from]
#[source]
urly::InvalidUrl,
),
}
impl From<std::convert::Infallible> for WsConfigError {
fn from(err: std::convert::Infallible) -> WsConfigError {
match err {}
}
}
#[derive(Debug, thiserror::Error)]
pub enum WsClientError {
#[error("Invalid client configuration: {0}")]
Config(
#[from]
#[source]
WsConfigError,
),
#[error("Invalid request")]
InvalidRequest(
#[from]
#[source]
EncodeError,
),
#[error("Invalid response")]
InvalidResponse(
#[from]
#[source]
DecodeError,
),
#[error("Invalid response status: {0}")]
InvalidResponseStatus(StatusCode),
#[error("Invalid upgrade header")]
InvalidUpgradeHeader,
#[error("Invalid connection header")]
InvalidConnectionHeader(HeaderValue),
#[error("Missing CONNECTION header")]
MissingConnectionHeader,
#[error("Missing SEC-WEBSOCKET-ACCEPT header")]
MissingWebSocketAcceptHeader,
#[error("Invalid challenge response")]
InvalidChallengeResponse(String, HeaderValue),
#[error("Invalid WebSocket subprotocol: {0:?}")]
InvalidWebSocketProtocol(HeaderValue),
#[error("Unexpected WebSocket extensions: {0:?}")]
UnexpectedWebSocketExtensions(HeaderValue),
#[error("{0}")]
Protocol(
#[from]
#[source]
ProtocolError,
),
#[error("Timeout while waiting for response")]
Timeout,
#[error("Failed to connect to host: {0}")]
Connect(
#[from]
#[source]
ConnectError,
),
#[error("Connector has been disconnected: {0:?}")]
Disconnected(#[source] Option<io::Error>),
}
impl From<Either<DecodeError, io::Error>> for WsClientError {
fn from(err: Either<DecodeError, io::Error>) -> Self {
match err {
Either::Left(err) => WsClientError::InvalidResponse(err),
Either::Right(err) => WsClientError::Disconnected(Some(err)),
}
}
}
impl From<Either<EncodeError, io::Error>> for WsClientError {
fn from(err: Either<EncodeError, io::Error>) -> Self {
match err {
Either::Left(err) => WsClientError::InvalidRequest(err),
Either::Right(err) => WsClientError::Disconnected(Some(err)),
}
}
}
impl Clone for WsClientError {
fn clone(&self) -> Self {
match self {
WsClientError::Config(err) => WsClientError::Config(err.clone()),
WsClientError::InvalidRequest(err) => WsClientError::InvalidRequest(err.clone()),
WsClientError::InvalidResponse(err) => WsClientError::InvalidResponse(*err),
WsClientError::InvalidResponseStatus(err) => WsClientError::InvalidResponseStatus(*err),
WsClientError::InvalidUpgradeHeader => WsClientError::InvalidUpgradeHeader,
WsClientError::InvalidConnectionHeader(err) => {
WsClientError::InvalidConnectionHeader(err.clone())
}
WsClientError::MissingConnectionHeader => WsClientError::MissingConnectionHeader,
WsClientError::MissingWebSocketAcceptHeader => {
WsClientError::MissingWebSocketAcceptHeader
}
WsClientError::InvalidChallengeResponse(n, val) => {
WsClientError::InvalidChallengeResponse(n.clone(), val.clone())
}
WsClientError::InvalidWebSocketProtocol(val) => {
WsClientError::InvalidWebSocketProtocol(val.clone())
}
WsClientError::UnexpectedWebSocketExtensions(val) => {
WsClientError::UnexpectedWebSocketExtensions(val.clone())
}
WsClientError::Protocol(err) => WsClientError::Protocol(*err),
WsClientError::Timeout => WsClientError::Timeout,
WsClientError::Connect(err) => WsClientError::Connect(err.clone()),
WsClientError::Disconnected(err) => {
WsClientError::Disconnected(err.as_ref().map(clone_io_error))
}
}
}
}
impl ErrorDiagnostic for WsClientError {
fn signature(&self) -> &'static str {
"ntex-ws-client"
}
}
#[derive(Copy, Clone, PartialEq, Eq, Debug, thiserror::Error)]
pub enum HandshakeError {
#[error("Method not allowed")]
GetMethodRequired,
#[error("Websocket upgrade is expected")]
NoWebsocketUpgrade,
#[error("Connection upgrade is expected")]
NoConnectionUpgrade,
#[error("Websocket version header is required")]
NoVersionHeader,
#[error("Unsupported version")]
UnsupportedVersion,
#[error("Unknown websocket key")]
BadWebsocketKey,
#[error("Invalid websocket subprotocol")]
BadWebsocketProtocol,
}
impl ResponseError for HandshakeError {
fn error_response(&self) -> Response {
match *self {
HandshakeError::GetMethodRequired => {
Response::MethodNotAllowed().header(ALLOW, "GET").build()
}
HandshakeError::NoWebsocketUpgrade => Response::BadRequest()
.reason("No WebSocket UPGRADE header found")
.build(),
HandshakeError::NoConnectionUpgrade => Response::BadRequest()
.reason("No CONNECTION upgrade")
.build(),
HandshakeError::NoVersionHeader => Response::BadRequest()
.reason("Websocket version header is required")
.header(SEC_WEBSOCKET_VERSION, "13")
.build(),
HandshakeError::UnsupportedVersion => Response::BadRequest()
.reason("Unsupported version")
.header(SEC_WEBSOCKET_VERSION, "13")
.build(),
HandshakeError::BadWebsocketKey => {
Response::BadRequest().reason("Handshake error").build()
}
HandshakeError::BadWebsocketProtocol => Response::BadRequest()
.reason("Invalid websocket subprotocol")
.build(),
}
}
}
impl ResponseError for ProtocolError {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_client_error_from_either() {
let err = WsClientError::from(Either::<DecodeError, io::Error>::Left(DecodeError::Method));
assert!(matches!(
err,
WsClientError::InvalidResponse(DecodeError::Method)
));
let err = WsClientError::from(Either::<DecodeError, _>::Right(io::Error::other("x")));
assert!(matches!(err, WsClientError::Disconnected(Some(_))));
let err = WsClientError::from(Either::<EncodeError, io::Error>::Left(
EncodeError::UnexpectedEof,
));
assert!(matches!(
err,
WsClientError::InvalidRequest(EncodeError::UnexpectedEof)
));
let err = WsClientError::from(Either::<EncodeError, _>::Right(io::Error::other("x")));
assert!(matches!(err, WsClientError::Disconnected(Some(_))));
}
#[test]
fn test_client_error_clone() {
let hdr = HeaderValue::from_static("v");
let errs = [
WsClientError::Config(WsConfigError::MissingHost),
WsClientError::InvalidRequest(EncodeError::UnexpectedEof),
WsClientError::InvalidResponse(DecodeError::Method),
WsClientError::InvalidResponseStatus(StatusCode::OK),
WsClientError::InvalidUpgradeHeader,
WsClientError::InvalidConnectionHeader(hdr.clone()),
WsClientError::MissingConnectionHeader,
WsClientError::MissingWebSocketAcceptHeader,
WsClientError::InvalidChallengeResponse("key".into(), hdr.clone()),
WsClientError::InvalidWebSocketProtocol(hdr.clone()),
WsClientError::UnexpectedWebSocketExtensions(hdr),
WsClientError::Protocol(ProtocolError::Overflow),
WsClientError::Timeout,
WsClientError::Connect(ConnectError::Unresolved),
WsClientError::Disconnected(None),
WsClientError::Disconnected(Some(io::Error::other("disconnected"))),
];
for err in errs {
assert_eq!(err.clone().to_string(), err.to_string());
assert_eq!(format!("{:?}", err.clone()), format!("{err:?}"));
assert_eq!(err.signature(), "ntex-ws-client");
}
}
}