1use 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
10pub fn handshake(req: &RequestHead) -> Result<ResponseBuilder, HandshakeError> {
17 verify_handshake(req)?;
18 Ok(handshake_response(req))
19}
20
21pub fn verify_handshake(req: &RequestHead) -> Result<(), HandshakeError> {
31 if req.method != Method::GET {
33 return Err(HandshakeError::GetMethodRequired);
34 }
35
36 if !header_contains_token(req.headers(), &header::UPGRADE, "websocket") {
38 return Err(HandshakeError::NoWebsocketUpgrade);
39 }
40
41 if !header_contains_token(req.headers(), &header::CONNECTION, "upgrade") {
43 return Err(HandshakeError::NoConnectionUpgrade);
44 }
45
46 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 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
82pub 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}