1use std::io::{Read, Write};
13
14use crate::ws_codec::apply_mask;
15
16const WEBSOCKET_GUID: &[u8] = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11";
18
19fn generate_client_nonce() -> [u8; 16] {
23 let mut state: u64 = 0x9E3779B97F4A7C15u64;
24 state ^= std::time::SystemTime::now()
25 .duration_since(std::time::UNIX_EPOCH)
26 .map(|d| d.as_nanos() as u64)
27 .unwrap_or(0xCAFEBABE);
28 let stack_addr = &state as *const _ as u64;
29 state ^= stack_addr;
30 let mut out = [0u8; 16];
31 for i in 0..16 {
32 state ^= state >> 12;
33 state ^= state << 25;
34 state ^= state >> 27;
35 out[i] = ((state.wrapping_mul(0x2545F4914F6CDD1D)) >> ((i % 8) * 8)) as u8;
36 }
37 out
38}
39
40pub fn generate_sec_websocket_key() -> String {
42 use bun_base64::encode_alloc;
43 let nonce = generate_client_nonce();
44 String::from_utf8(encode_alloc(&nonce)).unwrap_or_default()
45}
46
47fn read_line_direct<S: Read>(stream: &mut S) -> Result<String, HandshakeError> {
51 let mut buf = Vec::new();
52 let mut byte = [0u8; 1];
53 loop {
54 match stream.read(&mut byte) {
55 Ok(0) => return Err(HandshakeError::ReadError),
56 Ok(_) => {
57 buf.push(byte[0]);
58 if buf.ends_with(b"\r\n") {
59 buf.truncate(buf.len() - 2);
60 return String::from_utf8(buf).map_err(|_| HandshakeError::InvalidRequest);
61 }
62 if buf.len() > 8192 {
63 return Err(HandshakeError::InvalidRequest);
64 }
65 }
66 Err(_) => return Err(HandshakeError::ReadError),
67 }
68 }
69}
70
71pub fn server_handshake<S: Read + Write>(stream: &mut S) -> Result<(), HandshakeError> {
75 let request_line = read_line_direct(stream)?;
76
77 if !request_line.starts_with("GET ") {
78 return Err(HandshakeError::InvalidRequest);
79 }
80
81 let mut sec_websocket_key = None;
82 loop {
83 let line = read_line_direct(stream)?;
84 if line.is_empty() {
85 break;
86 }
87 if line.to_lowercase().starts_with("sec-websocket-key:") {
88 let key = line.split(':').nth(1).map(|s| s.trim());
89 if let Some(key) = key {
90 sec_websocket_key = Some(key.to_string());
91 }
92 }
93 }
94
95 let key = sec_websocket_key.ok_or(HandshakeError::MissingKey)?;
96 let accept = compute_accept(&key);
97
98 let response = format!(
99 "HTTP/1.1 101 Switching Protocols\r\n\
100 Upgrade: websocket\r\n\
101 Connection: Upgrade\r\n\
102 Sec-WebSocket-Accept: {}\r\n\
103 \r\n",
104 accept
105 );
106
107 stream
108 .write_all(response.as_bytes())
109 .map_err(|_| HandshakeError::WriteError)?;
110 stream.flush().map_err(|_| HandshakeError::WriteError)?;
111 Ok(())
112}
113
114pub fn client_handshake<S: Read + Write>(
120 stream: &mut S,
121 host: &str,
122 path: &str,
123) -> Result<String, HandshakeError> {
124 let sec_websocket_key = generate_sec_websocket_key();
125
126 let request = format!(
127 "GET {} HTTP/1.1\r\n\
128 Host: {}\r\n\
129 Upgrade: websocket\r\n\
130 Connection: Upgrade\r\n\
131 Sec-WebSocket-Key: {}\r\n\
132 Sec-WebSocket-Version: 13\r\n\
133 \r\n",
134 path, host, sec_websocket_key
135 );
136
137 stream
138 .write_all(request.as_bytes())
139 .map_err(|_| HandshakeError::WriteError)?;
140 stream
141 .flush()
142 .map_err(|_| HandshakeError::WriteError)?;
143
144 let status_line = read_line_direct(stream)?;
145 if !status_line.starts_with("HTTP/1.1 101") {
146 return Err(HandshakeError::InvalidRequest);
147 }
148
149 let mut accept_value: Option<String> = None;
150 loop {
151 let line = read_line_direct(stream)?;
152 if line.is_empty() {
153 break;
154 }
155 if line.to_lowercase().starts_with("sec-websocket-accept:") {
156 let val = line.split(':').nth(1).map(|s| s.trim());
157 if let Some(v) = val {
158 accept_value = Some(v.to_string());
159 }
160 }
161 }
162
163 let accept = accept_value.ok_or(HandshakeError::MissingKey)?;
164 let expected = compute_accept(&sec_websocket_key);
165 if accept != expected {
166 return Err(HandshakeError::MissingKey);
167 }
168
169 Ok(sec_websocket_key)
170}
171
172pub fn compute_accept(key: &str) -> String {
174 use bun_base64::encode_alloc;
175
176 let key_bytes = key.as_bytes();
177 let mut combined = Vec::with_capacity(key_bytes.len() + WEBSOCKET_GUID.len());
178 combined.extend_from_slice(key_bytes);
179 combined.extend_from_slice(WEBSOCKET_GUID);
180
181 let mut output = [0u8; 20];
184 unsafe {
185 bun_boringssl_sys::SHA1(combined.as_ptr(), combined.len(), output.as_mut_ptr());
186 }
187 String::from_utf8(encode_alloc(&output)).unwrap_or_default()
188}
189
190#[derive(Debug, Clone, PartialEq, Eq)]
191pub enum HandshakeError {
192 ReadError,
193 WriteError,
194 InvalidRequest,
195 MissingKey,
196}
197
198#[allow(dead_code)]
202fn _mask_link(payload: &mut [u8], key: &[u8; 4]) {
203 apply_mask(payload, key);
204}
205
206#[cfg(test)]
207mod tests {
208 use super::*;
209 use std::io::Cursor;
210
211 #[test]
212 fn compute_accept_known_vector() {
213 let accept = compute_accept("dGhlIHNhbXBsZSBub25jZQ==");
215 assert_eq!(accept, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=");
216 }
217
218 #[test]
219 fn server_handshake_valid_request() {
220 use std::io::Seek;
221 let request = b"GET /chat HTTP/1.1\r\n\
222 Host: example.com\r\n\
223 Upgrade: websocket\r\n\
224 Connection: Upgrade\r\n\
225 Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\
226 Sec-WebSocket-Version: 13\r\n\
227 \r\n";
228
229 let mut cursor = Cursor::new(request.to_vec());
230 let result = server_handshake(&mut cursor);
231 assert!(result.is_ok());
232
233 cursor.seek(std::io::SeekFrom::Start(0)).unwrap();
234 let mut output = Vec::new();
235 cursor.read_to_end(&mut output).unwrap();
236 let response = String::from_utf8(output).unwrap();
237 assert!(response.contains("HTTP/1.1 101 Switching Protocols"));
238 assert!(response.contains("Sec-WebSocket-Accept: s3pPLMBiTxaQ9kYGzzhZRbK+xOo="));
239 }
240
241 #[test]
242 fn server_handshake_missing_key() {
243 let request = b"GET /chat HTTP/1.1\r\nHost: example.com\r\n\r\n";
244 let mut cursor = Cursor::new(request.to_vec());
245 let result = server_handshake(&mut cursor);
246 assert!(matches!(result, Err(HandshakeError::MissingKey)));
247 }
248
249 #[test]
250 fn server_handshake_non_get_request() {
251 let request = b"POST /chat HTTP/1.1\r\nSec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\n\r\n";
252 let mut cursor = Cursor::new(request.to_vec());
253 let result = server_handshake(&mut cursor);
254 assert!(matches!(result, Err(HandshakeError::InvalidRequest)));
255 }
256
257 #[test]
258 fn generate_sec_websocket_key_is_base64_24_chars() {
259 let key = generate_sec_websocket_key();
260 assert_eq!(key.len(), 24);
261 assert!(
262 key.chars()
263 .all(|c| c.is_ascii_alphanumeric() || c == '+' || c == '/' || c == '='),
264 "invalid chars: {}",
265 key
266 );
267 }
268
269 #[test]
270 fn generate_sec_websocket_key_unique() {
271 let k1 = generate_sec_websocket_key();
272 let k2 = generate_sec_websocket_key();
273 assert_ne!(k1, k2);
274 }
275
276 #[test]
277 fn client_handshake_valid_response() {
278 let known_key = "dGhlIHNhbXBsZSBub25jZQ==";
279 let accept = compute_accept(known_key);
280 assert_eq!(accept, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=");
281
282 let mut buf: Vec<u8> = Vec::new();
283 let request = format!(
284 "GET / HTTP/1.1\r\nHost: test\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\
285 Sec-WebSocket-Key: {}\r\nSec-WebSocket-Version: 13\r\n\r\n",
286 known_key
287 );
288 buf.extend_from_slice(request.as_bytes());
289 let response = format!(
290 "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\
291 Sec-WebSocket-Accept: {}\r\n\r\n",
292 accept
293 );
294 buf.extend_from_slice(response.as_bytes());
295
296 let s = String::from_utf8(buf).unwrap();
297 assert!(s.contains("HTTP/1.1 101 Switching Protocols"));
298 assert!(s.contains("Sec-WebSocket-Accept: s3pPLMBiTxaQ9kYGzzhZRbK+xOo="));
299 }
300
301 #[test]
302 fn client_handshake_rejects_non_101_status() {
303 let mut buf: Vec<u8> = Vec::new();
304 buf.extend_from_slice(
305 b"GET / HTTP/1.1\r\nHost: test\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\
306 Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\nSec-WebSocket-Version: 13\r\n\r\n",
307 );
308 buf.extend_from_slice(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n");
309 let mut cursor = Cursor::new(buf);
310 let result = client_handshake(&mut cursor, "test", "/");
311 assert!(matches!(result, Err(HandshakeError::InvalidRequest)));
312 }
313
314 #[test]
315 fn client_handshake_rejects_wrong_accept() {
316 let mut buf: Vec<u8> = Vec::new();
317 buf.extend_from_slice(
318 b"GET / HTTP/1.1\r\nHost: test\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\
319 Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==\r\nSec-WebSocket-Version: 13\r\n\r\n",
320 );
321 buf.extend_from_slice(
322 b"HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\
323 Sec-WebSocket-Accept: aW52YWxpZGtleQ==\r\n\r\n",
324 );
325 let mut cursor = Cursor::new(buf);
326 let result = client_handshake(&mut cursor, "test", "/");
327 assert!(matches!(result, Err(HandshakeError::MissingKey)));
328 }
329}