Skip to main content

rama_http_headers/common/
sec_websocket_key.rs

1use base64::{Engine, engine::general_purpose::STANDARD};
2use rama_http_types::HeaderValue;
3
4use crate::{HeaderDecode, HeaderEncode, TypedHeader};
5
6/// The `Sec-WebSocket-Key` header.
7#[derive(Clone, Debug, PartialEq, Eq, Hash)]
8pub struct SecWebSocketKey(pub(super) HeaderValue);
9
10impl SecWebSocketKey {
11    #[must_use]
12    pub fn random() -> Self {
13        let r: [u8; 16] = rand::random();
14        r.into()
15    }
16}
17
18impl TypedHeader for SecWebSocketKey {
19    fn name() -> &'static ::rama_http_types::header::HeaderName {
20        &::rama_http_types::header::SEC_WEBSOCKET_KEY
21    }
22}
23
24impl HeaderDecode for SecWebSocketKey {
25    fn decode<'i, I>(values: &mut I) -> Result<Self, crate::Error>
26    where
27        I: Iterator<Item = &'i ::rama_http_types::header::HeaderValue>,
28    {
29        let value = crate::util::TryFromValues::try_from_values(values).map(SecWebSocketKey)?;
30        let mut k = [0u8; 16];
31        if STANDARD.decode_slice(value.0.as_bytes(), &mut k[..]).ok() != Some(16) {
32            Err(crate::Error::invalid())
33        } else {
34            Ok(value)
35        }
36    }
37}
38
39impl HeaderEncode for SecWebSocketKey {
40    fn encode<E: Extend<::rama_http_types::HeaderValue>>(&self, values: &mut E) {
41        values.extend(::std::iter::once((&self.0).into()));
42    }
43}
44
45impl From<[u8; 16]> for SecWebSocketKey {
46    fn from(bytes: [u8; 16]) -> Self {
47        #[expect(
48            clippy::expect_used,
49            reason = " ASSUMPTION standard base64 encoding of 16 bytes is always valid http header"
50        )]
51        let mut value = HeaderValue::try_from(STANDARD.encode(bytes))
52            .expect(" ASSUMPTION standard base64 encoding of 16 bytes is always valid http header");
53        value.set_sensitive(true);
54        Self(value)
55    }
56}
57
58#[cfg(test)]
59mod tests {
60    use super::*;
61    use crate::common::test_decode;
62
63    #[test]
64    fn from_bytes() {
65        let bytes: [u8; 16] = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16];
66        _ = SecWebSocketKey::from(bytes);
67    }
68
69    #[test]
70    fn test_invalid_websocket_key_empty() {
71        assert!(test_decode::<SecWebSocketKey>(&[""]).is_none());
72    }
73
74    #[test]
75    fn test_invalid_websocket_key_too_long() {
76        assert!(test_decode::<SecWebSocketKey>(&["dGhlIHNhbXBsZSBub25jZQ==AAAAAAAAAA"]).is_none());
77    }
78
79    #[test]
80    fn test_invalid_websocket_key_base64_symbol() {
81        assert!(test_decode::<SecWebSocketKey>(&["dGhlIHNhbXBsZSBub25jZQ!!"]).is_none());
82    }
83
84    #[test]
85    fn test_invalid_websocket_key_decoded_length() {
86        assert!(test_decode::<SecWebSocketKey>(&["AAAAAAAAAAAAAAAAAAAAAAAA"]).is_none());
87    }
88
89    #[test]
90    fn test_valid_websocket_key() {
91        _ = test_decode::<SecWebSocketKey>(&["dGhlIHNhbXBsZSBub25jZQ=="]).unwrap();
92    }
93}