1use std::fmt;
2
3use base64::{Engine, engine::general_purpose::STANDARD as base64};
4
5use super::error::HandshakeError;
6
7#[derive(Debug, Eq, PartialEq, Clone, Copy)]
9pub enum OpCode {
10 Continue,
12 Text,
14 Binary,
16 Close,
18 Ping,
20 Pong,
22}
23
24impl fmt::Display for OpCode {
25 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
26 match *self {
27 OpCode::Continue => write!(f, "CONTINUE"),
28 OpCode::Text => write!(f, "TEXT"),
29 OpCode::Binary => write!(f, "BINARY"),
30 OpCode::Close => write!(f, "CLOSE"),
31 OpCode::Ping => write!(f, "PING"),
32 OpCode::Pong => write!(f, "PONG"),
33 }
34 }
35}
36
37impl From<OpCode> for u8 {
38 fn from(code: OpCode) -> u8 {
39 match code {
40 OpCode::Continue => 0,
41 OpCode::Text => 1,
42 OpCode::Binary => 2,
43 OpCode::Close => 8,
44 OpCode::Ping => 9,
45 OpCode::Pong => 10,
46 }
47 }
48}
49
50impl TryFrom<u8> for OpCode {
51 type Error = ();
52
53 fn try_from(byte: u8) -> Result<OpCode, Self::Error> {
54 match byte {
55 0 => Ok(OpCode::Continue),
56 1 => Ok(OpCode::Text),
57 2 => Ok(OpCode::Binary),
58 8 => Ok(OpCode::Close),
59 9 => Ok(OpCode::Ping),
60 10 => Ok(OpCode::Pong),
61 _ => Err(()),
62 }
63 }
64}
65
66impl OpCode {
67 pub(crate) fn is_control(self) -> bool {
68 matches!(self, OpCode::Close | OpCode::Ping | OpCode::Pong)
69 }
70}
71
72#[derive(Debug, Eq, PartialEq, Clone, Copy)]
75pub enum CloseCode {
76 Normal,
79 Away,
82 Protocol,
85 Unsupported,
90 Abnormal,
95 Invalid,
100 Policy,
106 Size,
110 Extension,
118 Error,
122 Restart,
126 Again,
130 BadGateway,
132 Tls,
137 Other(u16),
141}
142
143impl From<CloseCode> for u16 {
144 fn from(code: CloseCode) -> u16 {
145 match code {
146 CloseCode::Normal => 1000,
147 CloseCode::Away => 1001,
148 CloseCode::Protocol => 1002,
149 CloseCode::Unsupported => 1003,
150 CloseCode::Abnormal => 1006,
151 CloseCode::Invalid => 1007,
152 CloseCode::Policy => 1008,
153 CloseCode::Size => 1009,
154 CloseCode::Extension => 1010,
155 CloseCode::Error => 1011,
156 CloseCode::Restart => 1012,
157 CloseCode::Again => 1013,
158 CloseCode::BadGateway => 1014,
159 CloseCode::Tls => 1015,
160 CloseCode::Other(code) => code,
161 }
162 }
163}
164
165impl From<u16> for CloseCode {
166 fn from(code: u16) -> CloseCode {
167 match code {
168 1000 => CloseCode::Normal,
169 1001 => CloseCode::Away,
170 1002 => CloseCode::Protocol,
171 1003 => CloseCode::Unsupported,
172 1006 => CloseCode::Abnormal,
173 1007 => CloseCode::Invalid,
174 1008 => CloseCode::Policy,
175 1009 => CloseCode::Size,
176 1010 => CloseCode::Extension,
177 1011 => CloseCode::Error,
178 1012 => CloseCode::Restart,
179 1013 => CloseCode::Again,
180 1014 => CloseCode::BadGateway,
181 1015 => CloseCode::Tls,
182 _ => CloseCode::Other(code),
183 }
184 }
185}
186
187impl CloseCode {
188 pub(crate) fn is_valid(self) -> bool {
189 !matches!(self, CloseCode::Abnormal | CloseCode::Tls)
190 && !matches!(self, CloseCode::Other(code) if !(3000..=4999).contains(&code))
191 }
192}
193
194#[derive(Debug, Eq, PartialEq, Clone)]
195pub struct CloseReason {
197 pub code: CloseCode,
199 pub description: Option<String>,
201}
202
203impl From<CloseCode> for CloseReason {
204 fn from(code: CloseCode) -> Self {
205 CloseReason {
206 code,
207 description: None,
208 }
209 }
210}
211
212impl<T: Into<String>> From<(CloseCode, T)> for CloseReason {
213 fn from(info: (CloseCode, T)) -> Self {
214 CloseReason {
215 code: info.0,
216 description: Some(info.1.into()),
217 }
218 }
219}
220
221const H0: u32 = 0x6745_2301;
223const H1: u32 = 0xEFCD_AB89;
224const H2: u32 = 0x98BA_DCFE;
225const H3: u32 = 0x1032_5476;
226const H4: u32 = 0xC3D2_E1F0;
227const WS_GUID: &str = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
228
229#[allow(clippy::many_single_char_names)]
230pub fn hash_key(key: &[u8]) -> Result<String, HandshakeError> {
237 if key.len() > 32 {
238 return Err(HandshakeError::BadWebsocketKey);
239 }
240 let mut input = [0; 192];
241 let klen = key.len();
242 let len = klen + 36;
243 input[..klen].copy_from_slice(key);
244 input[klen..len].copy_from_slice(WS_GUID.as_bytes());
245
246 let (mut h0, mut h1, mut h2, mut h3, mut h4) = (H0, H1, H2, H3, H4);
248 let (mut a, mut b, mut c, mut d, mut e);
249
250 let msg = pad_message(len, &mut input);
252
253 for chunk in msg.chunks(64) {
255 let mut schedule = [0u32; 80];
257 for (i, block) in chunk.chunks(4).enumerate() {
258 schedule[i] = u32::from_be_bytes(block.try_into().unwrap());
259 }
260 for i in 16..80 {
261 schedule[i] = schedule[i - 3] ^ schedule[i - 8] ^ schedule[i - 14] ^ schedule[i - 16];
262 schedule[i] = schedule[i].rotate_left(1);
263 }
264
265 a = h0;
266 b = h1;
267 c = h2;
268 d = h3;
269 e = h4;
270
271 for (i, sch) in schedule.iter().enumerate() {
273 let (f, k) = match i {
274 0..=19 => ((b & c) | ((!b) & d), 0x5A82_7999),
275 20..=39 => (b ^ c ^ d, 0x6ED9_EBA1),
276 40..=59 => ((b & c) | (b & d) | (c & d), 0x8F1B_BCDC),
277 _ => (b ^ c ^ d, 0xCA62_C1D6),
278 };
279
280 let temp = a
281 .rotate_left(5)
282 .wrapping_add(f)
283 .wrapping_add(e)
284 .wrapping_add(k)
285 .wrapping_add(*sch);
286 e = d;
287 d = c;
288 c = b.rotate_left(30);
289 b = a;
290 a = temp;
291 }
292
293 h0 = h0.wrapping_add(a);
295 h1 = h1.wrapping_add(b);
296 h2 = h2.wrapping_add(c);
297 h3 = h3.wrapping_add(d);
298 h4 = h4.wrapping_add(e);
299 }
300
301 let mut hash = [0u8; 20];
302 hash[0..4].copy_from_slice(&h0.to_be_bytes());
303 hash[4..8].copy_from_slice(&h1.to_be_bytes());
304 hash[8..12].copy_from_slice(&h2.to_be_bytes());
305 hash[12..16].copy_from_slice(&h3.to_be_bytes());
306 hash[16..20].copy_from_slice(&h4.to_be_bytes());
307
308 Ok(base64.encode(hash))
309}
310
311fn pad_message(len: usize, input: &mut [u8]) -> &[u8] {
312 let mut cur = len + 1;
313 let bit_length = len as u64 * 8;
314
315 input[len] = 0x80;
316 while (cur * 8) % 512 != 448 {
317 input[cur] = 0;
318 cur += 1;
319 }
320 let orig_len = &bit_length.to_be_bytes();
321 let total = cur + orig_len.len();
322 input[cur..total].copy_from_slice(orig_len);
323 &input[..total]
324}
325
326#[cfg(test)]
327#[allow(unused_imports, unused_variables, dead_code)]
328mod tests {
329 use super::*;
330
331 macro_rules! opcode_into {
332 ($from:expr => $opcode:pat) => {
333 match OpCode::try_from($from).unwrap() {
334 e @ $opcode => (),
335 e => unreachable!("{:?}", e),
336 }
337 };
338 }
339
340 macro_rules! opcode_from {
341 ($from:expr => $opcode:pat) => {
342 let res: u8 = $from.into();
343 match res {
344 e @ $opcode => (),
345 e => unreachable!("{:?}", e),
346 }
347 };
348 }
349
350 #[test]
351 fn test_to_opcode() {
352 opcode_into!(0 => OpCode::Continue);
353 opcode_into!(1 => OpCode::Text);
354 opcode_into!(2 => OpCode::Binary);
355 opcode_into!(8 => OpCode::Close);
356 opcode_into!(9 => OpCode::Ping);
357 opcode_into!(10 => OpCode::Pong);
358 assert!(OpCode::try_from(99).is_err());
359 }
360
361 #[test]
362 fn test_from_opcode() {
363 opcode_from!(OpCode::Continue => 0);
364 opcode_from!(OpCode::Text => 1);
365 opcode_from!(OpCode::Binary => 2);
366 opcode_from!(OpCode::Close => 8);
367 opcode_from!(OpCode::Ping => 9);
368 opcode_from!(OpCode::Pong => 10);
369 }
370
371 #[test]
372 fn test_from_opcode_display() {
373 assert_eq!(format!("{}", OpCode::Continue), "CONTINUE");
374 assert_eq!(format!("{}", OpCode::Text), "TEXT");
375 assert_eq!(format!("{}", OpCode::Binary), "BINARY");
376 assert_eq!(format!("{}", OpCode::Close), "CLOSE");
377 assert_eq!(format!("{}", OpCode::Ping), "PING");
378 assert_eq!(format!("{}", OpCode::Pong), "PONG");
379 }
380
381 #[test]
382 fn test_hash_key() {
383 let hash = hash_key(b"hello actix-web").unwrap();
384 assert_eq!(&hash, "cR1dlyUUJKp0s/Bel25u5TgvC3E=");
385 }
386
387 #[test]
388 fn closecode_from_u16() {
389 assert_eq!(CloseCode::from(1000u16), CloseCode::Normal);
390 assert_eq!(CloseCode::from(1001u16), CloseCode::Away);
391 assert_eq!(CloseCode::from(1002u16), CloseCode::Protocol);
392 assert_eq!(CloseCode::from(1003u16), CloseCode::Unsupported);
393 assert_eq!(CloseCode::from(1006u16), CloseCode::Abnormal);
394 assert_eq!(CloseCode::from(1007u16), CloseCode::Invalid);
395 assert_eq!(CloseCode::from(1008u16), CloseCode::Policy);
396 assert_eq!(CloseCode::from(1009u16), CloseCode::Size);
397 assert_eq!(CloseCode::from(1010u16), CloseCode::Extension);
398 assert_eq!(CloseCode::from(1011u16), CloseCode::Error);
399 assert_eq!(CloseCode::from(1012u16), CloseCode::Restart);
400 assert_eq!(CloseCode::from(1013u16), CloseCode::Again);
401 assert_eq!(CloseCode::from(1014u16), CloseCode::BadGateway);
402 assert_eq!(CloseCode::from(1015u16), CloseCode::Tls);
403 assert_eq!(CloseCode::from(3000u16), CloseCode::Other(3000));
404 assert!(!CloseCode::from(1005u16).is_valid());
405 assert!(!CloseCode::from(1006u16).is_valid());
406 assert!(!CloseCode::from(1015u16).is_valid());
407 assert!(!CloseCode::from(2000u16).is_valid());
408 }
409
410 #[test]
411 fn closecode_into_u16() {
412 assert_eq!(1000u16, Into::<u16>::into(CloseCode::Normal));
413 assert_eq!(1001u16, Into::<u16>::into(CloseCode::Away));
414 assert_eq!(1002u16, Into::<u16>::into(CloseCode::Protocol));
415 assert_eq!(1003u16, Into::<u16>::into(CloseCode::Unsupported));
416 assert_eq!(1006u16, Into::<u16>::into(CloseCode::Abnormal));
417 assert_eq!(1007u16, Into::<u16>::into(CloseCode::Invalid));
418 assert_eq!(1008u16, Into::<u16>::into(CloseCode::Policy));
419 assert_eq!(1009u16, Into::<u16>::into(CloseCode::Size));
420 assert_eq!(1010u16, Into::<u16>::into(CloseCode::Extension));
421 assert_eq!(1011u16, Into::<u16>::into(CloseCode::Error));
422 assert_eq!(1012u16, Into::<u16>::into(CloseCode::Restart));
423 assert_eq!(1013u16, Into::<u16>::into(CloseCode::Again));
424 assert_eq!(1014u16, Into::<u16>::into(CloseCode::BadGateway));
425 assert_eq!(1015u16, Into::<u16>::into(CloseCode::Tls));
426 assert_eq!(3000u16, Into::<u16>::into(CloseCode::Other(3000)));
427 assert!(!CloseCode::Other(2000).is_valid());
428 }
429
430 #[test]
431 fn test_hash_key_too_long() {
432 assert_eq!(hash_key(&[b'a'; 33]), Err(HandshakeError::BadWebsocketKey));
433 assert!(hash_key(&[b'a'; 32]).is_ok());
434 }
435}