Skip to main content

ntex/ws/
proto.rs

1use std::fmt;
2
3use base64::{Engine, engine::general_purpose::STANDARD as base64};
4
5use super::error::HandshakeError;
6
7/// WebSocket frame operation codes defined by RFC 6455.
8#[derive(Debug, Eq, PartialEq, Clone, Copy)]
9pub enum OpCode {
10    /// Indicates a continuation frame of a fragmented message.
11    Continue,
12    /// Indicates a text data frame.
13    Text,
14    /// Indicates a binary data frame.
15    Binary,
16    /// Indicates a close control frame.
17    Close,
18    /// Indicates a ping control frame.
19    Ping,
20    /// Indicates a pong control frame.
21    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/// Status code used to indicate why an endpoint is closing the `WebSocket`
73/// connection.
74#[derive(Debug, Eq, PartialEq, Clone, Copy)]
75pub enum CloseCode {
76    /// Indicates a normal closure, meaning that the purpose for
77    /// which the connection was established has been fulfilled.
78    Normal,
79    /// Indicates that an endpoint is "going away", such as a server
80    /// going down or a browser having navigated away from a page.
81    Away,
82    /// Indicates that an endpoint is terminating the connection due
83    /// to a protocol error.
84    Protocol,
85    /// Indicates that an endpoint is terminating the connection
86    /// because it has received a type of data it cannot accept (e.g., an
87    /// endpoint that understands only text data MAY send this if it
88    /// receives a binary message).
89    Unsupported,
90    /// Indicates that the connection closed without a close control frame.
91    ///
92    /// This code is reserved for local reporting and cannot be sent in a
93    /// close frame.
94    Abnormal,
95    /// Indicates that an endpoint is terminating the connection
96    /// because it has received data within a message that was not
97    /// consistent with the type of the message (e.g., non-UTF-8 \[RFC3629\]
98    /// data within a text message).
99    Invalid,
100    /// Indicates that an endpoint is terminating the connection
101    /// because it has received a message that violates its policy.  This
102    /// is a generic status code that can be returned when there is no
103    /// other more suitable status code (e.g., Unsupported or Size) or if there
104    /// is a need to hide specific details about the policy.
105    Policy,
106    /// Indicates that an endpoint is terminating the connection
107    /// because it has received a message that is too big for it to
108    /// process.
109    Size,
110    /// Indicates that an endpoint (client) is terminating the
111    /// connection because it has expected the server to negotiate one or
112    /// more extension, but the server didn't return them in the response
113    /// message of the WebSocket handshake.  The list of extensions that
114    /// are needed should be given as the reason for closing.
115    /// Note that this status code is not used by the server, because it
116    /// can fail the WebSocket handshake instead.
117    Extension,
118    /// Indicates that a server is terminating the connection because
119    /// it encountered an unexpected condition that prevented it from
120    /// fulfilling the request.
121    Error,
122    /// Indicates that the server is restarting. A client may choose to
123    /// reconnect, and if it does, it should use a randomized delay of 5-30
124    /// seconds between attempts.
125    Restart,
126    /// Indicates that the server is overloaded and the client should either
127    /// connect to a different IP (when multiple targets exist), or
128    /// reconnect to the same IP when a user has performed an action.
129    Again,
130    /// Indicates that an upstream server returned an invalid response.
131    BadGateway,
132    /// Indicates that the TLS handshake failed.
133    ///
134    /// This code is reserved for local reporting and cannot be sent in a
135    /// close frame.
136    Tls,
137    /// An unrecognized or application-defined close code.
138    ///
139    /// Only values in the range 3000 through 4999 can be sent.
140    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)]
195/// Reason supplied in a WebSocket close control frame.
196pub struct CloseReason {
197    /// Close status code.
198    pub code: CloseCode,
199    /// Optional human-readable description.
200    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
221// SHA-1 hashing algorithm initial hash values.
222const 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)]
230/// Computes the `Sec-WebSocket-Accept` value for a client handshake key.
231///
232/// # Errors
233///
234/// Returns [`HandshakeError::BadWebsocketKey`] when `key` is longer than
235/// 32 bytes.
236pub 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    // Initialize variables to the SHA-1's initial hash values.
247    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    // Pad the key
251    let msg = pad_message(len, &mut input);
252
253    // Process each 512-bit chunk of the padded message.
254    for chunk in msg.chunks(64) {
255        // Get the message schedule
256        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        // Main loop of the SHA-1 algorithm
272        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        // Add the compressed chunk to the current hash value.
294        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}