1extern crate alloc;
37
38use core::ops::Deref;
39
40use http::{
41 Method, StatusCode, Version,
42 header::{
43 ALLOW, CONNECTION, HeaderMap, HeaderValue, SEC_WEBSOCKET_ACCEPT, SEC_WEBSOCKET_KEY, SEC_WEBSOCKET_VERSION,
44 UPGRADE,
45 },
46 request::Request,
47 response::{Builder, Response},
48 uri::Uri,
49};
50
51mod codec;
52mod crypto;
53mod error;
54mod frame;
55mod mask;
56mod proto;
57
58pub use self::{
59 codec::{Codec, Item, Message},
60 error::{HandshakeError, ProtocolError},
61 proto::{CloseCode, CloseReason, OpCode, hash_key},
62};
63
64#[allow(clippy::declare_interior_mutable_const)]
65mod const_header {
66 use super::HeaderValue;
67
68 pub(super) const WEBSOCKET: HeaderValue = HeaderValue::from_static("websocket");
69 pub(super) const UPGRADE_VALUE: HeaderValue = HeaderValue::from_static("upgrade");
70 pub(super) const SEC_WEBSOCKET_VERSION_VALUE: HeaderValue = HeaderValue::from_static("13");
71}
72
73use const_header::*;
74
75impl From<HandshakeError> for Builder {
76 fn from(e: HandshakeError) -> Self {
77 match e {
78 HandshakeError::GetMethodRequired => Response::builder()
79 .status(StatusCode::METHOD_NOT_ALLOWED)
80 .header(ALLOW, "GET"),
81
82 _ => Response::builder().status(StatusCode::BAD_REQUEST),
83 }
84 }
85}
86
87pub fn client_request_from_uri(uri: Uri, version: Version) -> Request<()> {
92 let mut req = Request::new(());
93 *req.uri_mut() = uri;
94 *req.version_mut() = version;
95
96 client_request_extend(&mut req);
97
98 req
99}
100
101pub fn client_request_extend<B>(req: &mut Request<B>) {
111 match req.version() {
112 Version::HTTP_11 => {
113 req.headers_mut().insert(UPGRADE, WEBSOCKET);
114 req.headers_mut().insert(CONNECTION, UPGRADE_VALUE);
115
116 let output = crypto::base64::<16, 24>(&crypto::random());
118
119 req.headers_mut()
120 .insert(SEC_WEBSOCKET_KEY, HeaderValue::from_bytes(&output).unwrap());
121 }
122 Version::HTTP_2 => {
123 *req.method_mut() = Method::CONNECT;
124 req.extensions_mut().insert(Http2WsProtocol::new());
125 }
126 _ => {}
127 }
128
129 req.headers_mut()
130 .insert(SEC_WEBSOCKET_VERSION, SEC_WEBSOCKET_VERSION_VALUE);
131}
132
133#[derive(Clone)]
134pub struct Http2WsProtocol(&'static str);
135
136impl AsRef<str> for Http2WsProtocol {
137 fn as_ref(&self) -> &str {
138 self.0
139 }
140}
141
142impl Deref for Http2WsProtocol {
143 type Target = str;
144
145 fn deref(&self) -> &Self::Target {
146 self.0
147 }
148}
149
150impl Http2WsProtocol {
151 const fn new() -> Self {
152 Self("websocket")
153 }
154}
155
156pub fn handshake(method: &Method, headers: &HeaderMap) -> Result<Builder, HandshakeError> {
158 let key = verify_handshake(method, headers)?;
159 let builder = handshake_response(key);
160 Ok(builder)
161}
162
163pub fn handshake_h2(method: &Method, headers: &HeaderMap) -> Result<Builder, HandshakeError> {
173 if method != Method::CONNECT {
175 return Err(HandshakeError::ConnectMethodRequired);
176 }
177
178 ws_version_check(headers)?;
179
180 Ok(Response::builder().status(StatusCode::OK))
181}
182
183fn verify_handshake<'a>(method: &'a Method, headers: &'a HeaderMap) -> Result<&'a [u8], HandshakeError> {
185 if method != Method::GET {
187 return Err(HandshakeError::GetMethodRequired);
188 }
189
190 let has_upgrade_hd = headers
192 .get(UPGRADE)
193 .and_then(|hdr| hdr.to_str().ok())
194 .filter(|s| s.to_ascii_lowercase().contains("websocket"))
195 .is_some();
196
197 if !has_upgrade_hd {
198 return Err(HandshakeError::NoWebsocketUpgrade);
199 }
200
201 let has_connection_hd = headers
203 .get(CONNECTION)
204 .and_then(|hdr| hdr.to_str().ok())
205 .filter(|s| s.to_ascii_lowercase().contains("upgrade"))
206 .is_some();
207
208 if !has_connection_hd {
209 return Err(HandshakeError::NoConnectionUpgrade);
210 }
211
212 ws_version_check(headers)?;
213
214 let value = headers.get(SEC_WEBSOCKET_KEY).ok_or(HandshakeError::BadWebsocketKey)?;
216
217 Ok(value.as_bytes())
218}
219
220fn handshake_response(key: &[u8]) -> Builder {
224 let key = hash_key(key);
225
226 Response::builder()
227 .status(StatusCode::SWITCHING_PROTOCOLS)
228 .header(UPGRADE, WEBSOCKET)
229 .header(CONNECTION, UPGRADE_VALUE)
230 .header(
231 SEC_WEBSOCKET_ACCEPT,
232 HeaderValue::from_bytes(&key).unwrap(),
234 )
235}
236
237fn ws_version_check(headers: &HeaderMap) -> Result<(), HandshakeError> {
239 let value = headers
240 .get(SEC_WEBSOCKET_VERSION)
241 .ok_or(HandshakeError::NoVersionHeader)?;
242
243 if value != "13" && value != "8" && value != "7" {
244 Err(HandshakeError::UnsupportedVersion)
245 } else {
246 Ok(())
247 }
248}
249
250#[cfg(feature = "stream")]
251pub mod stream;
252
253#[cfg(feature = "stream")]
254pub use self::stream::{RequestStream, ResponseSender, ResponseStream, ResponseWeakSender, WsError};
255
256#[cfg(feature = "stream")]
257pub type WsOutput<B> = (RequestStream<B>, Response<ResponseStream>, ResponseSender);
258
259#[cfg(feature = "stream")]
260pub fn ws<ReqB, B, T, E>(req: &Request<ReqB>, body: B) -> Result<WsOutput<B>, HandshakeError>
313where
314 B: futures_core::Stream<Item = Result<T, E>>,
315 T: AsRef<[u8]>,
316{
317 let builder = match req.version() {
318 Version::HTTP_2 => handshake_h2(req.method(), req.headers())?,
319 _ => handshake(req.method(), req.headers())?,
320 };
321
322 let decode = RequestStream::new(body);
323 let (res, tx) = decode.response_stream();
324
325 let res = builder
326 .body(res)
327 .expect("handshake function failed to generate correct Response Builder");
328
329 Ok((decode, res, tx))
330}
331
332#[cfg(test)]
333mod tests {
334 use super::*;
335
336 #[test]
337 fn test_handshake() {
338 let req = Request::builder().method(Method::POST).body(()).unwrap();
339 assert_eq!(
340 HandshakeError::GetMethodRequired,
341 verify_handshake(req.method(), req.headers()).unwrap_err(),
342 );
343
344 let req = Request::builder().body(()).unwrap();
345 assert_eq!(
346 HandshakeError::NoWebsocketUpgrade,
347 verify_handshake(req.method(), req.headers()).unwrap_err(),
348 );
349
350 let req = Request::builder()
351 .header(UPGRADE, HeaderValue::from_static("test"))
352 .body(())
353 .unwrap();
354 assert_eq!(
355 HandshakeError::NoWebsocketUpgrade,
356 verify_handshake(req.method(), req.headers()).unwrap_err(),
357 );
358
359 let req = Request::builder().header(UPGRADE, WEBSOCKET).body(()).unwrap();
360 assert_eq!(
361 HandshakeError::NoConnectionUpgrade,
362 verify_handshake(req.method(), req.headers()).unwrap_err(),
363 );
364
365 let req = Request::builder()
366 .header(UPGRADE, WEBSOCKET)
367 .header(CONNECTION, UPGRADE_VALUE)
368 .body(())
369 .unwrap();
370 assert_eq!(
371 HandshakeError::NoVersionHeader,
372 verify_handshake(req.method(), req.headers()).unwrap_err(),
373 );
374
375 let req = Request::builder()
376 .header(UPGRADE, WEBSOCKET)
377 .header(CONNECTION, UPGRADE_VALUE)
378 .header(SEC_WEBSOCKET_VERSION, HeaderValue::from_static("5"))
379 .body(())
380 .unwrap();
381 assert_eq!(
382 HandshakeError::UnsupportedVersion,
383 verify_handshake(req.method(), req.headers()).unwrap_err(),
384 );
385
386 let builder = || {
387 Request::builder()
388 .header(UPGRADE, WEBSOCKET)
389 .header(CONNECTION, UPGRADE_VALUE)
390 .header(SEC_WEBSOCKET_VERSION, SEC_WEBSOCKET_VERSION_VALUE)
391 };
392
393 let req = builder().body(()).unwrap();
394 assert_eq!(
395 HandshakeError::BadWebsocketKey,
396 verify_handshake(req.method(), req.headers()).unwrap_err(),
397 );
398
399 let req = builder()
400 .header(SEC_WEBSOCKET_KEY, SEC_WEBSOCKET_VERSION_VALUE)
401 .body(())
402 .unwrap();
403 let key = verify_handshake(req.method(), req.headers()).unwrap();
404 assert_eq!(
405 StatusCode::SWITCHING_PROTOCOLS,
406 handshake_response(key).body(()).unwrap().status()
407 );
408 }
409
410 #[test]
411 fn test_ws_error_http_response() {
412 let res = Builder::from(HandshakeError::GetMethodRequired).body(()).unwrap();
413 assert_eq!(res.status(), StatusCode::METHOD_NOT_ALLOWED);
414 let res = Builder::from(HandshakeError::NoWebsocketUpgrade).body(()).unwrap();
415 assert_eq!(res.status(), StatusCode::BAD_REQUEST);
416 let res = Builder::from(HandshakeError::NoConnectionUpgrade).body(()).unwrap();
417 assert_eq!(res.status(), StatusCode::BAD_REQUEST);
418 let res = Builder::from(HandshakeError::NoVersionHeader).body(()).unwrap();
419 assert_eq!(res.status(), StatusCode::BAD_REQUEST);
420 let res = Builder::from(HandshakeError::UnsupportedVersion).body(()).unwrap();
421 assert_eq!(res.status(), StatusCode::BAD_REQUEST);
422 let res = Builder::from(HandshakeError::BadWebsocketKey).body(()).unwrap();
423 assert_eq!(res.status(), StatusCode::BAD_REQUEST);
424 }
425}