use std::fmt::{self, Debug, Formatter};
use std::pin::Pin;
use std::task::{Context, Poll, ready};
use futures_util::sink::{Sink, SinkExt};
use futures_util::stream::{Stream, StreamExt};
use futures_util::{FutureExt, TryFutureExt, future};
use hyper::upgrade::OnUpgrade;
use salvo_core::http::header::{ORIGIN, SEC_WEBSOCKET_PROTOCOL, SEC_WEBSOCKET_VERSION, UPGRADE};
use salvo_core::http::headers::{
Connection, HeaderMapExt, SecWebsocketAccept, SecWebsocketKey, Upgrade,
};
use salvo_core::http::{StatusCode, StatusError};
use salvo_core::rt::tokio::TokioIo;
use salvo_core::conn::ConnCtrl;
use salvo_core::{Error, Request, Response};
use tokio_tungstenite::WebSocketStream;
use tokio_tungstenite::tungstenite::Bytes;
use tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode;
use tokio_tungstenite::tungstenite::protocol::frame::{CloseFrame, Utf8Bytes};
use tokio_tungstenite::tungstenite::protocol::{self, WebSocketConfig};
#[allow(missing_debug_implementations)]
pub struct WebSocketUpgrade {
config: Option<WebSocketConfig>,
protocols: Vec<String>,
accept_any: bool,
origin_check: Option<OriginCheck>,
}
enum OriginCheck {
Allowlist(Vec<String>),
Predicate(Box<dyn Fn(Option<&str>) -> bool + Send + Sync>),
}
impl OriginCheck {
fn allows(&self, origin: Option<&str>) -> bool {
match self {
Self::Allowlist(list) => {
origin.is_some_and(|o| list.iter().any(|allowed| allowed.eq_ignore_ascii_case(o)))
}
Self::Predicate(predicate) => predicate(origin),
}
}
}
impl Default for WebSocketUpgrade {
#[inline]
fn default() -> Self {
Self::new()
}
}
impl WebSocketUpgrade {
#[inline]
#[must_use]
pub fn new() -> Self {
Self {
config: None,
protocols: Vec::new(),
accept_any: false,
origin_check: None,
}
}
#[inline]
#[must_use]
pub fn with_config(config: WebSocketConfig) -> Self {
Self {
config: Some(config),
protocols: Vec::new(),
accept_any: false,
origin_check: None,
}
}
#[inline]
#[must_use]
pub fn protocols(mut self, protocols: &[&str]) -> Self {
self.protocols = protocols.iter().map(|s| (*s).to_owned()).collect();
self
}
#[inline]
#[must_use]
pub fn accept_any_protocol(mut self) -> Self {
self.accept_any = true;
self
}
#[inline]
#[must_use]
pub fn allowed_origins(mut self, origins: &[&str]) -> Self {
self.origin_check = Some(OriginCheck::Allowlist(
origins.iter().map(|s| (*s).to_owned()).collect(),
));
self
}
#[inline]
#[must_use]
pub fn check_origin(
mut self,
predicate: impl Fn(Option<&str>) -> bool + Send + Sync + 'static,
) -> Self {
self.origin_check = Some(OriginCheck::Predicate(Box::new(predicate)));
self
}
#[inline]
#[must_use]
pub fn write_buffer_size(mut self, max: usize) -> Self {
self.config
.get_or_insert_with(WebSocketConfig::default)
.write_buffer_size = max;
self
}
#[inline]
#[must_use]
pub fn max_write_buffer_size(mut self, max: usize) -> Self {
self.config
.get_or_insert_with(WebSocketConfig::default)
.max_write_buffer_size = max;
self
}
#[inline]
#[must_use]
pub fn max_message_size(mut self, max: usize) -> Self {
self.config
.get_or_insert_with(WebSocketConfig::default)
.max_message_size = Some(max);
self
}
#[inline]
#[must_use]
pub fn max_frame_size(mut self, max: usize) -> Self {
self.config
.get_or_insert_with(WebSocketConfig::default)
.max_frame_size = Some(max);
self
}
#[inline]
#[must_use]
pub fn accept_unmasked_frames(mut self, accept: bool) -> Self {
self.config
.get_or_insert_with(WebSocketConfig::default)
.accept_unmasked_frames = accept;
self
}
pub async fn upgrade<F, Fut>(
&self,
req: &mut Request,
res: &mut Response,
callback: F,
) -> Result<(), StatusError>
where
F: FnOnce(WebSocket) -> Fut + Send + 'static,
Fut: Future<Output = ()> + Send + 'static,
{
let req_headers = req.headers();
let matched = req_headers
.typed_get::<Connection>()
.map(|conn| conn.contains(UPGRADE))
.unwrap_or(false);
if !matched {
tracing::debug!("missing connection upgrade");
return Err(StatusError::bad_request().brief("missing connection upgrade"));
}
let matched = req_headers
.get(UPGRADE)
.and_then(|v| v.to_str().ok())
.is_some_and(|v| v.eq_ignore_ascii_case("websocket"));
if !matched {
tracing::debug!("missing upgrade header or it is not websocket");
return Err(StatusError::bad_request()
.brief("missing upgrade header or it is not websocket"));
}
let matched = !req_headers
.get(SEC_WEBSOCKET_VERSION)
.and_then(|v| v.to_str().ok())
.map(|v| v == "13")
.unwrap_or(false);
if matched {
tracing::debug!("websocket version is not 13");
return Err(StatusError::bad_request().brief("websocket version is not 13"));
}
let Some(sec_ws_key) = req_headers.typed_get::<SecWebsocketKey>() else {
tracing::debug!("missing sec-websocket-key header");
return Err(StatusError::bad_request().brief("missing sec-websocket-key header"));
};
if let Some(check) = &self.origin_check {
let origin = req_headers.get(ORIGIN).and_then(|v| v.to_str().ok());
if !check.allows(origin) {
tracing::debug!(?origin, "rejecting websocket upgrade: origin not allowed");
return Err(StatusError::forbidden().brief("websocket origin not allowed"));
}
}
res.status_code(StatusCode::SWITCHING_PROTOCOLS);
res.headers_mut().typed_insert(Connection::upgrade());
res.headers_mut().typed_insert(Upgrade::websocket());
res.headers_mut()
.typed_insert(SecWebsocketAccept::from(sec_ws_key));
let selected: Option<&str> = if self.accept_any {
find_first_client_protocol(req_headers)
} else {
select_protocol(req_headers, &self.protocols)
};
if let Some(selected) = selected {
match selected.parse() {
Ok(value) => {
res.headers_mut().insert(SEC_WEBSOCKET_PROTOCOL, value);
}
Err(error) => {
tracing::debug!(
protocol = selected,
?error,
"skipping Sec-WebSocket-Protocol response header because the negotiated protocol is not a valid header value"
);
}
}
}
if let Some(on_upgrade) = req.extensions_mut().remove::<OnUpgrade>() {
if let Some(conn_ctrl) = req.extensions().get::<ConnCtrl>() {
conn_ctrl.relax_timeouts();
}
let config = self.config;
tokio::spawn(async move {
let socket = on_upgrade
.and_then(move |upgraded| {
tracing::debug!("websocket upgrade complete");
WebSocket::from_raw_socket(upgraded, protocol::Role::Server, config).map(Ok)
})
.await;
match socket {
Ok(socket) => callback(socket).await,
Err(error) => {
tracing::debug!(?error, "websocket connection upgrade failed");
}
}
});
Ok(())
} else {
tracing::debug!("websocket cannot be upgraded because no upgrade state was present");
Err(StatusError::bad_request()
.brief("websocket cannot be upgraded because no upgrade state was present"))
}
}
}
pub struct WebSocket {
inner: WebSocketStream<TokioIo<hyper::upgrade::Upgraded>>,
}
impl WebSocket {
#[inline]
pub(crate) async fn from_raw_socket(
upgraded: hyper::upgrade::Upgraded,
role: protocol::Role,
config: Option<protocol::WebSocketConfig>,
) -> Self {
WebSocketStream::from_raw_socket(TokioIo::new(upgraded), role, config)
.map(|inner| Self { inner })
.await
}
pub async fn recv(&mut self) -> Option<Result<Message, Error>> {
self.next().await
}
pub async fn send(&mut self, msg: Message) -> Result<(), Error> {
self.inner.send(msg.inner).await.map_err(Error::other)
}
#[inline]
pub async fn close(mut self) -> Result<(), Error> {
future::poll_fn(|cx| Pin::new(&mut self).poll_close(cx)).await
}
}
impl Stream for WebSocket {
type Item = Result<Message, Error>;
#[inline]
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Option<Self::Item>> {
match ready!(Pin::new(&mut self.inner).poll_next(cx)) {
Some(Ok(item)) => Poll::Ready(Some(Ok(Message { inner: item }))),
Some(Err(e)) => {
tracing::debug!("websocket poll error: {}", e);
Poll::Ready(Some(Err(Error::other(e))))
}
None => {
tracing::debug!("websocket closed");
Poll::Ready(None)
}
}
}
}
impl Sink<Message> for WebSocket {
type Error = Error;
#[inline]
fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
Pin::new(&mut self.inner)
.poll_ready(cx)
.map_err(Error::other)
}
#[inline]
fn start_send(mut self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
Pin::new(&mut self.inner)
.start_send(item.inner)
.map_err(Error::other)
}
#[inline]
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
Pin::new(&mut self.inner)
.poll_flush(cx)
.map_err(Error::other)
}
#[inline]
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<Result<(), Self::Error>> {
Pin::new(&mut self.inner)
.poll_close(cx)
.map_err(Error::other)
}
}
impl Debug for WebSocket {
#[inline]
fn fmt(&self, f: &mut Formatter) -> fmt::Result {
f.debug_struct("WebSocket").finish()
}
}
#[derive(Eq, PartialEq, Clone)]
pub struct Message {
inner: protocol::Message,
}
impl Message {
#[inline]
pub fn text<S: Into<Utf8Bytes>>(s: S) -> Self {
Self {
inner: protocol::Message::text(s),
}
}
#[inline]
pub fn binary<V: Into<Bytes>>(v: V) -> Self {
Self {
inner: protocol::Message::binary(v),
}
}
#[inline]
pub fn ping<V: Into<Bytes>>(v: V) -> Self {
Self {
inner: protocol::Message::Ping(v.into()),
}
}
#[inline]
pub fn pong<V: Into<Bytes>>(v: V) -> Self {
Self {
inner: protocol::Message::Pong(v.into()),
}
}
#[inline]
pub fn close() -> Self {
Self {
inner: protocol::Message::Close(None),
}
}
#[inline]
pub fn close_with(code: impl Into<u16>, reason: impl Into<Utf8Bytes>) -> Self {
Self {
inner: protocol::Message::Close(Some(CloseFrame {
code: CloseCode::from(code.into()),
reason: reason.into(),
})),
}
}
#[inline]
pub fn is_text(&self) -> bool {
self.inner.is_text()
}
#[inline]
pub fn is_binary(&self) -> bool {
self.inner.is_binary()
}
#[inline]
pub fn is_close(&self) -> bool {
self.inner.is_close()
}
#[inline]
pub fn is_ping(&self) -> bool {
self.inner.is_ping()
}
#[inline]
pub fn is_pong(&self) -> bool {
self.inner.is_pong()
}
#[inline]
pub fn close_frame(&self) -> Option<(u16, &str)> {
if let protocol::Message::Close(Some(ref close_frame)) = self.inner {
Some((close_frame.code.into(), close_frame.reason.as_ref()))
} else {
None
}
}
#[inline]
pub fn as_str(&self) -> Result<&str, Error> {
match &self.inner {
protocol::Message::Text(s) => Ok(s.as_str()),
_ => Err(Error::Other("not a text message".into())),
}
}
#[inline]
pub fn as_bytes(&self) -> &[u8] {
match &self.inner {
protocol::Message::Text(s) => s.as_bytes(),
protocol::Message::Binary(v)
| protocol::Message::Ping(v)
| protocol::Message::Pong(v) => v.as_ref(),
protocol::Message::Close(_) => &[],
protocol::Message::Frame(v) => v.payload(),
}
}
}
impl Debug for Message {
#[inline]
fn fmt(&self, f: &mut Formatter) -> fmt::Result {
fmt::Debug::fmt(&self.inner, f)
}
}
#[allow(clippy::from_over_into)]
impl Into<Vec<u8>> for Message {
#[inline]
fn into(self) -> Vec<u8> {
self.as_bytes().into()
}
}
fn find_first_client_protocol(req_headers: &salvo_core::http::HeaderMap) -> Option<&str> {
for value in req_headers.get_all(SEC_WEBSOCKET_PROTOCOL) {
let Ok(value) = value.to_str() else {
continue;
};
for token in value.split(',').map(str::trim) {
if !token.is_empty() {
return Some(token);
}
}
}
None
}
fn select_protocol<'a>(
req_headers: &'a salvo_core::http::HeaderMap,
protocols: &[String],
) -> Option<&'a str> {
if protocols.is_empty() {
return None;
}
for value in req_headers.get_all(SEC_WEBSOCKET_PROTOCOL) {
let Ok(value) = value.to_str() else {
continue;
};
for requested in value.split(',').map(str::trim) {
if requested.is_empty() {
continue;
}
if protocols.iter().any(|supported| supported == requested) {
return Some(requested);
}
}
}
None
}
#[cfg(test)]
mod tests {
use salvo_core::conn::{Acceptor, Listener};
use salvo_core::http::header::*;
use salvo_core::prelude::*;
use salvo_core::test::{ResponseExt, TestClient};
use salvo_core::rt::tokio::TokioIo;
use super::*;
#[handler]
async fn connect(req: &mut Request, res: &mut Response) -> Result<(), StatusError> {
WebSocketUpgrade::new()
.upgrade(req, res, |mut ws| async move {
while let Some(msg) = ws.recv().await {
let Ok(msg) = msg else {
return;
};
if ws.send(msg).await.is_err() {
return;
}
}
})
.await
}
#[handler]
async fn connect_with_protocols(req: &mut Request, res: &mut Response) -> Result<(), StatusError> {
WebSocketUpgrade::new()
.protocols(&["chat.v1", "chat.v2"])
.upgrade(req, res, |mut ws| async move {
while let Some(msg) = ws.recv().await {
let Ok(msg) = msg else {
return;
};
if ws.send(msg).await.is_err() {
return;
}
}
})
.await
}
#[handler]
async fn connect_accept_any(req: &mut Request, res: &mut Response) -> Result<(), StatusError> {
WebSocketUpgrade::new()
.accept_any_protocol()
.upgrade(req, res, |mut ws| async move {
while let Some(msg) = ws.recv().await {
let Ok(msg) = msg else {
return;
};
if ws.send(msg).await.is_err() {
return;
}
}
})
.await
}
#[handler]
async fn connect_with_allowed_origins(
req: &mut Request,
res: &mut Response,
) -> Result<(), StatusError> {
WebSocketUpgrade::new()
.allowed_origins(&["https://allowed.example"])
.upgrade(req, res, |_ws| async move {})
.await
}
#[test]
fn origin_check_allowlist_and_predicate() {
let allow = OriginCheck::Allowlist(vec!["https://A.example".to_owned()]);
assert!(allow.allows(Some("https://a.example"))); assert!(!allow.allows(Some("https://evil.example")));
assert!(!allow.allows(None));
let pred = OriginCheck::Predicate(Box::new(|o| o.is_none() || o == Some("https://ok")));
assert!(pred.allows(None));
assert!(pred.allows(Some("https://ok")));
assert!(!pred.allows(Some("https://no")));
}
#[tokio::test]
async fn test_websocket_rejects_disallowed_origin() {
let router = Router::new().goal(connect_with_allowed_origins);
let mut response = TestClient::get("http://127.0.0.1:5801")
.add_header(CONNECTION, "Upgrade", true)
.add_header(UPGRADE, "websocket", true)
.add_header(SEC_WEBSOCKET_KEY, "6D69KGBOr4Re+Nj6zx9aQA==", true)
.add_header(SEC_WEBSOCKET_VERSION, "13", true)
.add_header(ORIGIN, "https://evil.example", true)
.send(router)
.await;
assert_eq!(response.status_code, Some(StatusCode::FORBIDDEN));
assert!(
response
.take_string()
.await
.unwrap()
.contains("origin not allowed")
);
}
#[tokio::test]
async fn test_websocket() {
let router = Router::new().goal(connect);
let acceptor = TcpListener::new("127.0.0.1:0").bind().await;
let addr = acceptor.holdings()[0]
.local_addr
.clone()
.into_std()
.unwrap();
tokio::spawn(async move {
Server::new(acceptor).serve(router).await;
});
let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(stream))
.await
.unwrap();
tokio::task::spawn(async move {
if let Err(err) = conn.await {
println!("Connection failed: {err:?}");
}
});
let req = hyper::Request::builder()
.uri(format!("http://{addr}"))
.header(UPGRADE, "websocket")
.header(CONNECTION, "Upgrade")
.header(SEC_WEBSOCKET_KEY, "6D69KGBOr4Re+Nj6zx9aQA==")
.header(SEC_WEBSOCKET_VERSION, "13")
.body(http_body_util::Empty::<hyper::body::Bytes>::new())
.unwrap();
let res = sender.send_request(req).await.unwrap();
assert_eq!(res.status(), StatusCode::SWITCHING_PROTOCOLS);
}
#[tokio::test]
async fn test_websocket_missing_connection_upgrade_returns_consistent_message() {
let router = Router::new().goal(connect);
let mut response = TestClient::get("http://127.0.0.1:5801")
.send(router)
.await;
assert_eq!(response.status_code, Some(StatusCode::BAD_REQUEST));
assert!(
response
.take_string()
.await
.unwrap()
.contains("missing connection upgrade")
);
}
#[tokio::test]
async fn test_websocket_missing_sec_websocket_key_returns_consistent_message() {
let router = Router::new().goal(connect);
let mut response = TestClient::get("http://127.0.0.1:5801")
.add_header(UPGRADE, "websocket", true)
.add_header(CONNECTION, "Upgrade", true)
.add_header(SEC_WEBSOCKET_VERSION, "13", true)
.send(router)
.await;
assert_eq!(response.status_code, Some(StatusCode::BAD_REQUEST));
assert!(
response
.take_string()
.await
.unwrap()
.contains("missing sec-websocket-key header")
);
}
#[tokio::test]
async fn test_websocket_with_protocol_match() {
let router = Router::new().goal(connect_with_protocols);
let acceptor = TcpListener::new("127.0.0.1:0").bind().await;
let addr = acceptor.holdings()[0]
.local_addr
.clone()
.into_std()
.unwrap();
tokio::spawn(async move {
Server::new(acceptor).serve(router).await;
});
let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(stream))
.await
.unwrap();
tokio::task::spawn(async move {
if let Err(err) = conn.await {
println!("Connection failed: {err:?}");
}
});
let req = hyper::Request::builder()
.uri(format!("http://{addr}"))
.header(UPGRADE, "websocket")
.header(CONNECTION, "Upgrade")
.header(SEC_WEBSOCKET_KEY, "6D69KGBOr4Re+Nj6zx9aQA==")
.header(SEC_WEBSOCKET_VERSION, "13")
.header(SEC_WEBSOCKET_PROTOCOL, "chat.v2, chat.v1")
.body(http_body_util::Empty::<hyper::body::Bytes>::new())
.unwrap();
let res = sender.send_request(req).await.unwrap();
assert_eq!(res.status(), StatusCode::SWITCHING_PROTOCOLS);
assert_eq!(
res.headers().get(SEC_WEBSOCKET_PROTOCOL),
Some(&HeaderValue::from_static("chat.v2")),
);
}
#[tokio::test]
async fn test_websocket_with_protocol_no_match() {
let router = Router::new().goal(connect_with_protocols);
let acceptor = TcpListener::new("127.0.0.1:0").bind().await;
let addr = acceptor.holdings()[0]
.local_addr
.clone()
.into_std()
.unwrap();
tokio::spawn(async move {
Server::new(acceptor).serve(router).await;
});
let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(stream))
.await
.unwrap();
tokio::task::spawn(async move {
if let Err(err) = conn.await {
println!("Connection failed: {err:?}");
}
});
let req = hyper::Request::builder()
.uri(format!("http://{addr}"))
.header(UPGRADE, "websocket")
.header(CONNECTION, "Upgrade")
.header(SEC_WEBSOCKET_KEY, "6D69KGBOr4Re+Nj6zx9aQA==")
.header(SEC_WEBSOCKET_VERSION, "13")
.header(SEC_WEBSOCKET_PROTOCOL, "unknown-protocol")
.body(http_body_util::Empty::<hyper::body::Bytes>::new())
.unwrap();
let res = sender.send_request(req).await.unwrap();
assert_eq!(res.status(), StatusCode::SWITCHING_PROTOCOLS);
assert!(res.headers().get(SEC_WEBSOCKET_PROTOCOL).is_none());
}
#[tokio::test]
async fn test_websocket_without_protocol() {
let router = Router::new().goal(connect_with_protocols);
let acceptor = TcpListener::new("127.0.0.1:0").bind().await;
let addr = acceptor.holdings()[0]
.local_addr
.clone()
.into_std()
.unwrap();
tokio::spawn(async move {
Server::new(acceptor).serve(router).await;
});
let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(stream))
.await
.unwrap();
tokio::task::spawn(async move {
if let Err(err) = conn.await {
println!("Connection failed: {err:?}");
}
});
let req = hyper::Request::builder()
.uri(format!("http://{addr}"))
.header(UPGRADE, "websocket")
.header(CONNECTION, "Upgrade")
.header(SEC_WEBSOCKET_KEY, "6D69KGBOr4Re+Nj6zx9aQA==")
.header(SEC_WEBSOCKET_VERSION, "13")
.body(http_body_util::Empty::<hyper::body::Bytes>::new())
.unwrap();
let res = sender.send_request(req).await.unwrap();
assert_eq!(res.status(), StatusCode::SWITCHING_PROTOCOLS);
assert!(res.headers().get(SEC_WEBSOCKET_PROTOCOL).is_none());
}
#[tokio::test]
async fn test_websocket_accept_any_protocol() {
let router = Router::new().goal(connect_accept_any);
let acceptor = TcpListener::new("127.0.0.1:0").bind().await;
let addr = acceptor.holdings()[0]
.local_addr
.clone()
.into_std()
.unwrap();
tokio::spawn(async move {
Server::new(acceptor).serve(router).await;
});
let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(stream))
.await
.unwrap();
tokio::task::spawn(async move {
if let Err(err) = conn.await {
println!("Connection failed: {err:?}");
}
});
let req = hyper::Request::builder()
.uri(format!("http://{addr}"))
.header(UPGRADE, "websocket")
.header(CONNECTION, "Upgrade")
.header(SEC_WEBSOCKET_KEY, "6D69KGBOr4Re+Nj6zx9aQA==")
.header(SEC_WEBSOCKET_VERSION, "13")
.header(SEC_WEBSOCKET_PROTOCOL, "any-custom-protocol, fallback")
.body(http_body_util::Empty::<hyper::body::Bytes>::new())
.unwrap();
let res = sender.send_request(req).await.unwrap();
assert_eq!(res.status(), StatusCode::SWITCHING_PROTOCOLS);
assert_eq!(
res.headers().get(SEC_WEBSOCKET_PROTOCOL),
Some(&HeaderValue::from_static("any-custom-protocol")),
);
}
#[tokio::test]
async fn test_websocket_accept_any_protocol_leading_comma() {
let router = Router::new().goal(connect_accept_any);
let acceptor = TcpListener::new("127.0.0.1:0").bind().await;
let addr = acceptor.holdings()[0]
.local_addr
.clone()
.into_std()
.unwrap();
tokio::spawn(async move {
Server::new(acceptor).serve(router).await;
});
let stream = tokio::net::TcpStream::connect(addr).await.unwrap();
let (mut sender, conn) = hyper::client::conn::http1::handshake(TokioIo::new(stream))
.await
.unwrap();
tokio::task::spawn(async move {
if let Err(err) = conn.await {
println!("Connection failed: {err:?}");
}
});
let req = hyper::Request::builder()
.uri(format!("http://{addr}"))
.header(UPGRADE, "websocket")
.header(CONNECTION, "Upgrade")
.header(SEC_WEBSOCKET_KEY, "6D69KGBOr4Re+Nj6zx9aQA==")
.header(SEC_WEBSOCKET_VERSION, "13")
.header(SEC_WEBSOCKET_PROTOCOL, ", chat.v1")
.body(http_body_util::Empty::<hyper::body::Bytes>::new())
.unwrap();
let res = sender.send_request(req).await.unwrap();
assert_eq!(res.status(), StatusCode::SWITCHING_PROTOCOLS);
assert_eq!(
res.headers().get(SEC_WEBSOCKET_PROTOCOL),
Some(&HeaderValue::from_static("chat.v1")),
);
}
}