use crate::handshake::{UPGRADE_STR, WEBSOCKET_STR, WEBSOCKET_VERSION_STR};
use crate::test_fixture::{mock, ReadError};
use crate::{
accept_with, Error, ErrorKind, HttpError, NoExtProvider, ProtocolRegistry, WebSocketConfig,
};
use bitflags::_core::convert::Infallible;
use bytes::BytesMut;
use http::header::HeaderName;
use http::{HeaderMap, HeaderValue, Request, Response, Version};
use httparse::Header;
use ratchet_ext::{
Extension, ExtensionDecoder, ExtensionEncoder, ExtensionProvider, FrameHeader,
ReunitableExtension, RsvBits, SplittableExtension,
};
impl From<ReadError<httparse::Error>> for Error {
fn from(e: ReadError<httparse::Error>) -> Self {
Error::with_cause(ErrorKind::Http, e)
}
}
async fn exec_request(request: Request<()>) -> Result<Response<()>, Error> {
let (mut client, server) = mock();
client.write_request(request).await?;
let upgrader = accept_with(
server,
WebSocketConfig::default(),
NoExtProvider,
ProtocolRegistry::default(),
)
.await?;
let _upgraded = upgrader.upgrade().await?;
client.read_response().await.map_err(Into::into)
}
#[tokio::test]
async fn valid_response() {
let response = exec_request(valid_request()).await.unwrap();
let expected = Response::builder()
.status(101)
.header(http::header::CONNECTION, "upgrade")
.header(http::header::UPGRADE, WEBSOCKET_STR)
.header(
http::header::SEC_WEBSOCKET_ACCEPT,
"s3pPLMBiTxaQ9kYGzzhZRbK+xOo=",
)
.version(Version::HTTP_11)
.body(())
.unwrap();
assert_response_eq(response, expected);
}
#[tokio::test]
async fn bad_request() {
async fn t(request: Request<()>, name: HeaderName) {
match exec_request(request).await {
Ok(o) => panic!("Expected a test failure. Got: {:?}", o),
Err(e) => match e.downcast_ref::<HttpError>() {
Some(err) => {
assert_eq!(err, &HttpError::MissingHeader(name))
}
None => {
panic!("Expected a HTTP error. Got: {:?}", e)
}
},
}
}
t(
Request::builder()
.uri("/test")
.header(http::header::UPGRADE, WEBSOCKET_STR)
.header(http::header::SEC_WEBSOCKET_VERSION, WEBSOCKET_VERSION_STR)
.header(http::header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.header(http::header::HOST, "localtoast")
.body(())
.unwrap(),
http::header::CONNECTION,
)
.await;
t(
Request::builder()
.uri("/test")
.header(http::header::CONNECTION, UPGRADE_STR)
.header(http::header::SEC_WEBSOCKET_VERSION, WEBSOCKET_VERSION_STR)
.header(http::header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.header(http::header::HOST, "localtoast")
.body(())
.unwrap(),
http::header::UPGRADE,
)
.await;
t(
Request::builder()
.uri("/test")
.header(http::header::CONNECTION, UPGRADE_STR)
.header(http::header::UPGRADE, WEBSOCKET_STR)
.header(http::header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.header(http::header::HOST, "localtoast")
.body(())
.unwrap(),
http::header::SEC_WEBSOCKET_VERSION,
)
.await;
t(
Request::builder()
.uri("/test")
.header(http::header::CONNECTION, UPGRADE_STR)
.header(http::header::UPGRADE, WEBSOCKET_STR)
.header(http::header::SEC_WEBSOCKET_VERSION, WEBSOCKET_VERSION_STR)
.header(http::header::HOST, "localtoast")
.body(())
.unwrap(),
http::header::SEC_WEBSOCKET_KEY,
)
.await;
t(
Request::builder()
.uri("/test")
.header(http::header::CONNECTION, UPGRADE_STR)
.header(http::header::UPGRADE, WEBSOCKET_STR)
.header(http::header::SEC_WEBSOCKET_VERSION, WEBSOCKET_VERSION_STR)
.header(http::header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.body(())
.unwrap(),
http::header::HOST,
)
.await;
let request = Request::builder()
.uri("/test")
.version(Version::HTTP_10)
.header(http::header::CONNECTION, UPGRADE_STR)
.header(http::header::UPGRADE, WEBSOCKET_STR)
.header(http::header::SEC_WEBSOCKET_VERSION, WEBSOCKET_VERSION_STR)
.header(http::header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.header(http::header::HOST, "localtoast")
.body(())
.unwrap();
match exec_request(request).await {
Ok(o) => panic!("Expected a test failure. Got: {:?}", o),
Err(e) => match e.downcast_ref::<HttpError>() {
Some(err) => {
assert_eq!(err, &HttpError::HttpVersion(Some(0)));
}
None => {
panic!("Expected a HTTP error. Got: {:?}", e)
}
},
}
let request = Request::builder()
.uri("/test")
.method("post")
.header(http::header::CONNECTION, UPGRADE_STR)
.header(http::header::UPGRADE, WEBSOCKET_STR)
.header(http::header::SEC_WEBSOCKET_VERSION, WEBSOCKET_VERSION_STR)
.header(http::header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.header(http::header::HOST, "localtoast")
.body(())
.unwrap();
match exec_request(request).await {
Ok(o) => panic!("Expected a test failure. Got: {:?}", o),
Err(e) => match e.downcast_ref::<HttpError>() {
Some(err) => {
assert_eq!(err, &HttpError::HttpMethod(Some("post".to_string())));
}
None => {
panic!("Expected a HTTP error. Got: {:?}", e)
}
},
}
}
fn assert_response_eq(left: Response<()>, right: Response<()>) {
assert!(left.version().eq(&right.version()));
assert!(left.status().eq(&right.status()));
assert!(left.headers().eq(right.headers()));
}
#[derive(Debug, thiserror::Error)]
#[error("Error")]
struct ExtErr;
impl From<ExtErr> for Error {
fn from(e: ExtErr) -> Self {
Error::with_cause(ErrorKind::Extension, e)
}
}
struct BadExtProvider;
impl ExtensionProvider for BadExtProvider {
type Extension = Ext;
type Error = ExtErr;
fn apply_headers(&self, _headers: &mut HeaderMap) {}
fn negotiate_client(
&self,
_headers: &[Header],
) -> Result<Option<Self::Extension>, Self::Error> {
panic!("Unexpected client negotitation request")
}
fn negotiate_server(
&self,
_headers: &[Header],
) -> Result<Option<(Self::Extension, HeaderValue)>, Self::Error> {
Err(ExtErr)
}
}
#[derive(Copy, Clone, Debug)]
struct Ext;
impl ExtensionEncoder for Ext {
type Error = Infallible;
fn encode(
&mut self,
_payload: &mut BytesMut,
_header: &mut FrameHeader,
) -> Result<(), Self::Error> {
Ok(())
}
}
impl ExtensionDecoder for Ext {
type Error = Infallible;
fn decode(
&mut self,
_payload: &mut BytesMut,
_header: &mut FrameHeader,
) -> Result<(), Self::Error> {
Ok(())
}
}
impl Extension for Ext {
fn bits(&self) -> RsvBits {
RsvBits {
rsv1: true,
rsv2: false,
rsv3: false,
}
}
}
impl SplittableExtension for Ext {
type SplitEncoder = Self;
type SplitDecoder = Self;
fn split(self) -> (Self::SplitEncoder, Self::SplitDecoder) {
(self, self)
}
}
impl ReunitableExtension for Ext {
fn reunite(encoder: Self::SplitEncoder, _decoder: Self::SplitDecoder) -> Self {
encoder
}
}
fn valid_request() -> Request<()> {
Request::builder()
.uri("/test")
.header(http::header::CONNECTION, UPGRADE_STR)
.header(http::header::UPGRADE, WEBSOCKET_STR)
.header(http::header::SEC_WEBSOCKET_VERSION, WEBSOCKET_VERSION_STR)
.header(http::header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==")
.header(http::header::HOST, "localtoast")
.body(())
.unwrap()
}
#[tokio::test]
async fn bad_extension() {
let (mut client, server) = mock();
client.write_request(valid_request()).await.unwrap();
let result = accept_with(
server,
WebSocketConfig::default(),
BadExtProvider,
ProtocolRegistry::default(),
)
.await;
match result {
Ok(_) => {
panic!("Expected the connection to fail")
}
Err(e) => {
if e.downcast_ref::<ExtErr>().is_none() {
panic!("{:?}", e);
}
}
}
}