use base64::{Engine, engine::general_purpose::STANDARD as base64};
use crate::http::header::HeaderName;
use crate::http::{HeaderMap, RequestHead, Response, ResponseBuilder};
use crate::http::{Method, StatusCode, header};
use super::error::HandshakeError;
pub fn handshake(req: &RequestHead) -> Result<ResponseBuilder, HandshakeError> {
verify_handshake(req)?;
Ok(handshake_response(req))
}
pub fn verify_handshake(req: &RequestHead) -> Result<(), HandshakeError> {
if req.method != Method::GET {
return Err(HandshakeError::GetMethodRequired);
}
if !header_contains_token(req.headers(), &header::UPGRADE, "websocket") {
return Err(HandshakeError::NoWebsocketUpgrade);
}
if !header_contains_token(req.headers(), &header::CONNECTION, "upgrade") {
return Err(HandshakeError::NoConnectionUpgrade);
}
if !req.headers().contains_key(header::SEC_WEBSOCKET_VERSION) {
return Err(HandshakeError::NoVersionHeader);
}
let mut versions = req.headers().get_all(header::SEC_WEBSOCKET_VERSION);
if versions.next().is_none_or(|ver| ver != "13") || versions.next().is_some() {
return Err(HandshakeError::UnsupportedVersion);
}
let mut keys = req.headers().get_all(header::SEC_WEBSOCKET_KEY);
let valid_key = keys
.next()
.and_then(|key| base64.decode(key.as_bytes()).ok())
.is_some_and(|key| key.len() == 16)
&& keys.next().is_none();
if !valid_key {
return Err(HandshakeError::BadWebsocketKey);
}
Ok(())
}
pub(super) fn header_contains_token(
headers: &HeaderMap,
name: &HeaderName,
expected: &str,
) -> bool {
headers.get_all(name).any(|value| {
value.to_str().is_ok_and(|value| {
value
.split(',')
.any(|token| token.trim().eq_ignore_ascii_case(expected))
})
})
}
pub fn handshake_response(req: &RequestHead) -> ResponseBuilder {
let key = {
let key = req.headers().get(header::SEC_WEBSOCKET_KEY).unwrap();
crate::ws::hash_key(key.as_ref()).expect("validated Sec-WebSocket-Key")
};
Response::builder(StatusCode::SWITCHING_PROTOCOLS)
.upgrade("websocket")
.header(header::SEC_WEBSOCKET_ACCEPT, key)
.take()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::http::{error::ResponseError, test::TestRequest};
#[test]
fn test_handshake() {
let req = TestRequest::default().method(Method::POST).build();
assert_eq!(
HandshakeError::GetMethodRequired,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default().build();
assert_eq!(
HandshakeError::NoWebsocketUpgrade,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default()
.header(header::UPGRADE, header::HeaderValue::from_static("test"))
.build();
assert_eq!(
HandshakeError::NoWebsocketUpgrade,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default()
.header(
header::UPGRADE,
header::HeaderValue::from_static("notwebsocket"),
)
.build();
assert_eq!(
HandshakeError::NoWebsocketUpgrade,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default()
.header(
header::UPGRADE,
header::HeaderValue::from_static("WebSocket"),
)
.build();
assert_eq!(
HandshakeError::NoConnectionUpgrade,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default()
.header(
header::UPGRADE,
header::HeaderValue::from_static("websocket"),
)
.header(
header::CONNECTION,
header::HeaderValue::from_static("keep-alive, Upgrade"),
)
.build();
assert_eq!(
HandshakeError::NoVersionHeader,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default()
.header(
header::UPGRADE,
header::HeaderValue::from_static("websocket"),
)
.header(
header::CONNECTION,
header::HeaderValue::from_static("keep-alive, upgraded"),
)
.build();
assert_eq!(
HandshakeError::NoConnectionUpgrade,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default()
.header(
header::UPGRADE,
header::HeaderValue::from_static("websocket"),
)
.header(
header::CONNECTION,
header::HeaderValue::from_static("upgrade"),
)
.header(
header::SEC_WEBSOCKET_VERSION,
header::HeaderValue::from_static("5"),
)
.build();
assert_eq!(
HandshakeError::UnsupportedVersion,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default()
.header(
header::UPGRADE,
header::HeaderValue::from_static("websocket"),
)
.header(
header::CONNECTION,
header::HeaderValue::from_static("upgrade"),
)
.header(
header::SEC_WEBSOCKET_VERSION,
header::HeaderValue::from_static("13"),
)
.build();
assert_eq!(
HandshakeError::BadWebsocketKey,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default()
.header(
header::UPGRADE,
header::HeaderValue::from_static("websocket"),
)
.header(
header::CONNECTION,
header::HeaderValue::from_static("upgrade"),
)
.header(
header::SEC_WEBSOCKET_VERSION,
header::HeaderValue::from_static("13"),
)
.header(
header::SEC_WEBSOCKET_KEY,
header::HeaderValue::from_static("13"),
)
.build();
assert_eq!(
HandshakeError::BadWebsocketKey,
verify_handshake(req.head()).err().unwrap()
);
let req = TestRequest::default()
.header(
header::UPGRADE,
header::HeaderValue::from_static("websocket"),
)
.header(
header::CONNECTION,
header::HeaderValue::from_static("upgrade"),
)
.header(
header::SEC_WEBSOCKET_VERSION,
header::HeaderValue::from_static("13"),
)
.header(
header::SEC_WEBSOCKET_KEY,
header::HeaderValue::from_static("dGhlIHNhbXBsZSBub25jZQ=="),
)
.build();
verify_handshake(req.head()).unwrap();
let response = handshake_response(req.head()).build();
assert_eq!(StatusCode::SWITCHING_PROTOCOLS, response.status());
assert!(!response.headers().contains_key(header::TRANSFER_ENCODING));
}
#[test]
fn test_only_version_13_is_supported() {
let req = |versions: &[&'static str]| {
let mut req = TestRequest::default();
req.header(header::UPGRADE, "websocket")
.header(header::CONNECTION, "upgrade")
.header(header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==");
for ver in versions {
req.header(header::SEC_WEBSOCKET_VERSION, *ver);
}
req.build()
};
assert!(verify_handshake(req(&["13"]).head()).is_ok());
for versions in [&["8"][..], &["7"], &["13", "8"]] {
assert_eq!(
verify_handshake(req(versions).head()),
Err(HandshakeError::UnsupportedVersion)
);
}
}
#[test]
fn test_wserror_http_response() {
let resp: Response = HandshakeError::GetMethodRequired.error_response();
assert_eq!(resp.status(), StatusCode::METHOD_NOT_ALLOWED);
let resp: Response = HandshakeError::NoWebsocketUpgrade.error_response();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
let resp: Response = HandshakeError::NoConnectionUpgrade.error_response();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
let resp: Response = HandshakeError::NoVersionHeader.error_response();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
let resp: Response = HandshakeError::UnsupportedVersion.error_response();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
assert_eq!(
resp.headers().get(header::SEC_WEBSOCKET_VERSION).unwrap(),
"13"
);
let resp: Response = HandshakeError::BadWebsocketKey.error_response();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
let resp: Response = HandshakeError::BadWebsocketProtocol.error_response();
assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
}
}