Skip to main content

ntex/ws/
handshake.rs

1//! WebSocket opening-handshake helpers.
2use base64::{Engine, engine::general_purpose::STANDARD as base64};
3
4use crate::http::header::HeaderName;
5use crate::http::{HeaderMap, RequestHead, Response, ResponseBuilder};
6use crate::http::{Method, StatusCode, header};
7
8use super::error::HandshakeError;
9
10/// Verifies a WebSocket opening-handshake request and creates its response.
11///
12/// # Errors
13///
14/// Returns [`HandshakeError`] when the request method or required upgrade
15/// headers are invalid.
16pub fn handshake(req: &RequestHead) -> Result<ResponseBuilder, HandshakeError> {
17    verify_handshake(req)?;
18    Ok(handshake_response(req))
19}
20
21/// Verifies a WebSocket opening-handshake request.
22///
23/// The request must use `GET`, request a connection upgrade to WebSocket,
24/// include a `Sec-WebSocket-Key`, and use WebSocket version 13.
25///
26/// # Errors
27///
28/// Returns [`HandshakeError`] when the request method or required upgrade
29/// headers are invalid.
30pub fn verify_handshake(req: &RequestHead) -> Result<(), HandshakeError> {
31    // WebSocket accepts only GET
32    if req.method != Method::GET {
33        return Err(HandshakeError::GetMethodRequired);
34    }
35
36    // Check for "UPGRADE" to websocket header
37    if !header_contains_token(req.headers(), &header::UPGRADE, "websocket") {
38        return Err(HandshakeError::NoWebsocketUpgrade);
39    }
40
41    // Upgrade connection
42    if !header_contains_token(req.headers(), &header::CONNECTION, "upgrade") {
43        return Err(HandshakeError::NoConnectionUpgrade);
44    }
45
46    // check supported version
47    if !req.headers().contains_key(header::SEC_WEBSOCKET_VERSION) {
48        return Err(HandshakeError::NoVersionHeader);
49    }
50    let mut versions = req.headers().get_all(header::SEC_WEBSOCKET_VERSION);
51    if versions.next().is_none_or(|ver| ver != "13") || versions.next().is_some() {
52        return Err(HandshakeError::UnsupportedVersion);
53    }
54
55    // check client handshake for validity
56    let mut keys = req.headers().get_all(header::SEC_WEBSOCKET_KEY);
57    let valid_key = keys
58        .next()
59        .and_then(|key| base64.decode(key.as_bytes()).ok())
60        .is_some_and(|key| key.len() == 16)
61        && keys.next().is_none();
62    if !valid_key {
63        return Err(HandshakeError::BadWebsocketKey);
64    }
65    Ok(())
66}
67
68pub(super) fn header_contains_token(
69    headers: &HeaderMap,
70    name: &HeaderName,
71    expected: &str,
72) -> bool {
73    headers.get_all(name).any(|value| {
74        value.to_str().is_ok_and(|value| {
75            value
76                .split(',')
77                .any(|token| token.trim().eq_ignore_ascii_case(expected))
78        })
79    })
80}
81
82/// Creates a WebSocket opening-handshake response.
83///
84/// The returned response builder has status `101 Switching Protocols` and the
85/// required upgrade and challenge-response headers.
86///
87/// # Panics
88///
89/// Panics if `req` does not contain a `Sec-WebSocket-Key` header or the key
90/// exceeds the length accepted by [`hash_key`](crate::ws::hash_key). Use
91/// [`handshake`] when the request has not already been validated.
92pub fn handshake_response(req: &RequestHead) -> ResponseBuilder {
93    let key = {
94        let key = req.headers().get(header::SEC_WEBSOCKET_KEY).unwrap();
95        crate::ws::hash_key(key.as_ref()).expect("validated Sec-WebSocket-Key")
96    };
97
98    Response::builder(StatusCode::SWITCHING_PROTOCOLS)
99        .upgrade("websocket")
100        .header(header::SEC_WEBSOCKET_ACCEPT, key)
101        .take()
102}
103
104#[cfg(test)]
105mod tests {
106    use super::*;
107    use crate::http::{error::ResponseError, test::TestRequest};
108
109    #[test]
110    fn test_handshake() {
111        let req = TestRequest::default().method(Method::POST).build();
112        assert_eq!(
113            HandshakeError::GetMethodRequired,
114            verify_handshake(req.head()).err().unwrap()
115        );
116
117        let req = TestRequest::default().build();
118        assert_eq!(
119            HandshakeError::NoWebsocketUpgrade,
120            verify_handshake(req.head()).err().unwrap()
121        );
122
123        let req = TestRequest::default()
124            .header(header::UPGRADE, header::HeaderValue::from_static("test"))
125            .build();
126        assert_eq!(
127            HandshakeError::NoWebsocketUpgrade,
128            verify_handshake(req.head()).err().unwrap()
129        );
130
131        let req = TestRequest::default()
132            .header(
133                header::UPGRADE,
134                header::HeaderValue::from_static("notwebsocket"),
135            )
136            .build();
137        assert_eq!(
138            HandshakeError::NoWebsocketUpgrade,
139            verify_handshake(req.head()).err().unwrap()
140        );
141
142        let req = TestRequest::default()
143            .header(
144                header::UPGRADE,
145                header::HeaderValue::from_static("WebSocket"),
146            )
147            .build();
148        assert_eq!(
149            HandshakeError::NoConnectionUpgrade,
150            verify_handshake(req.head()).err().unwrap()
151        );
152
153        let req = TestRequest::default()
154            .header(
155                header::UPGRADE,
156                header::HeaderValue::from_static("websocket"),
157            )
158            .header(
159                header::CONNECTION,
160                header::HeaderValue::from_static("keep-alive, Upgrade"),
161            )
162            .build();
163        assert_eq!(
164            HandshakeError::NoVersionHeader,
165            verify_handshake(req.head()).err().unwrap()
166        );
167
168        let req = TestRequest::default()
169            .header(
170                header::UPGRADE,
171                header::HeaderValue::from_static("websocket"),
172            )
173            .header(
174                header::CONNECTION,
175                header::HeaderValue::from_static("keep-alive, upgraded"),
176            )
177            .build();
178        assert_eq!(
179            HandshakeError::NoConnectionUpgrade,
180            verify_handshake(req.head()).err().unwrap()
181        );
182
183        let req = TestRequest::default()
184            .header(
185                header::UPGRADE,
186                header::HeaderValue::from_static("websocket"),
187            )
188            .header(
189                header::CONNECTION,
190                header::HeaderValue::from_static("upgrade"),
191            )
192            .header(
193                header::SEC_WEBSOCKET_VERSION,
194                header::HeaderValue::from_static("5"),
195            )
196            .build();
197        assert_eq!(
198            HandshakeError::UnsupportedVersion,
199            verify_handshake(req.head()).err().unwrap()
200        );
201
202        let req = TestRequest::default()
203            .header(
204                header::UPGRADE,
205                header::HeaderValue::from_static("websocket"),
206            )
207            .header(
208                header::CONNECTION,
209                header::HeaderValue::from_static("upgrade"),
210            )
211            .header(
212                header::SEC_WEBSOCKET_VERSION,
213                header::HeaderValue::from_static("13"),
214            )
215            .build();
216        assert_eq!(
217            HandshakeError::BadWebsocketKey,
218            verify_handshake(req.head()).err().unwrap()
219        );
220
221        let req = TestRequest::default()
222            .header(
223                header::UPGRADE,
224                header::HeaderValue::from_static("websocket"),
225            )
226            .header(
227                header::CONNECTION,
228                header::HeaderValue::from_static("upgrade"),
229            )
230            .header(
231                header::SEC_WEBSOCKET_VERSION,
232                header::HeaderValue::from_static("13"),
233            )
234            .header(
235                header::SEC_WEBSOCKET_KEY,
236                header::HeaderValue::from_static("13"),
237            )
238            .build();
239        assert_eq!(
240            HandshakeError::BadWebsocketKey,
241            verify_handshake(req.head()).err().unwrap()
242        );
243
244        let req = TestRequest::default()
245            .header(
246                header::UPGRADE,
247                header::HeaderValue::from_static("websocket"),
248            )
249            .header(
250                header::CONNECTION,
251                header::HeaderValue::from_static("upgrade"),
252            )
253            .header(
254                header::SEC_WEBSOCKET_VERSION,
255                header::HeaderValue::from_static("13"),
256            )
257            .header(
258                header::SEC_WEBSOCKET_KEY,
259                header::HeaderValue::from_static("dGhlIHNhbXBsZSBub25jZQ=="),
260            )
261            .build();
262        verify_handshake(req.head()).unwrap();
263        let response = handshake_response(req.head()).build();
264        assert_eq!(StatusCode::SWITCHING_PROTOCOLS, response.status());
265        assert!(!response.headers().contains_key(header::TRANSFER_ENCODING));
266    }
267
268    #[test]
269    fn test_only_version_13_is_supported() {
270        let req = |versions: &[&'static str]| {
271            let mut req = TestRequest::default();
272            req.header(header::UPGRADE, "websocket")
273                .header(header::CONNECTION, "upgrade")
274                .header(header::SEC_WEBSOCKET_KEY, "dGhlIHNhbXBsZSBub25jZQ==");
275            for ver in versions {
276                req.header(header::SEC_WEBSOCKET_VERSION, *ver);
277            }
278            req.build()
279        };
280        assert!(verify_handshake(req(&["13"]).head()).is_ok());
281        for versions in [&["8"][..], &["7"], &["13", "8"]] {
282            assert_eq!(
283                verify_handshake(req(versions).head()),
284                Err(HandshakeError::UnsupportedVersion)
285            );
286        }
287    }
288
289    #[test]
290    fn test_wserror_http_response() {
291        let resp: Response = HandshakeError::GetMethodRequired.error_response();
292        assert_eq!(resp.status(), StatusCode::METHOD_NOT_ALLOWED);
293        let resp: Response = HandshakeError::NoWebsocketUpgrade.error_response();
294        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
295        let resp: Response = HandshakeError::NoConnectionUpgrade.error_response();
296        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
297        let resp: Response = HandshakeError::NoVersionHeader.error_response();
298        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
299        let resp: Response = HandshakeError::UnsupportedVersion.error_response();
300        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
301        assert_eq!(
302            resp.headers().get(header::SEC_WEBSOCKET_VERSION).unwrap(),
303            "13"
304        );
305        let resp: Response = HandshakeError::BadWebsocketKey.error_response();
306        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
307        let resp: Response = HandshakeError::BadWebsocketProtocol.error_response();
308        assert_eq!(resp.status(), StatusCode::BAD_REQUEST);
309    }
310}