/src/suricata8/rust/src/websocket/websocket.rs
Line | Count | Source |
1 | | /* Copyright (C) 2023-2025 Open Information Security Foundation |
2 | | * |
3 | | * You can copy, redistribute or modify this Program under the terms of |
4 | | * the GNU General Public License version 2 as published by the Free |
5 | | * Software Foundation. |
6 | | * |
7 | | * This program is distributed in the hope that it will be useful, |
8 | | * but WITHOUT ANY WARRANTY; without even the implied warranty of |
9 | | * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the |
10 | | * GNU General Public License for more details. |
11 | | * |
12 | | * You should have received a copy of the GNU General Public License |
13 | | * version 2 along with this program; if not, write to the Free Software |
14 | | * Foundation, Inc., 51 Franklin Street, Fifth Floor, Boston, MA |
15 | | * 02110-1301, USA. |
16 | | */ |
17 | | |
18 | | use super::parser; |
19 | | use crate::applayer::{self, *}; |
20 | | use crate::conf::conf_get; |
21 | | use crate::core::{ |
22 | | sc_app_layer_parser_trigger_raw_stream_inspection, ALPROTO_FAILED, ALPROTO_UNKNOWN, IPPROTO_TCP, |
23 | | }; |
24 | | use crate::direction::Direction; |
25 | | use crate::flow::Flow; |
26 | | use crate::frames::Frame; |
27 | | |
28 | | use nom7 as nom; |
29 | | use nom7::Needed; |
30 | | |
31 | | use flate2::Decompress; |
32 | | use flate2::FlushDecompress; |
33 | | use suricata_sys::sys::{ |
34 | | AppLayerParserState, AppProto, SCAppLayerParserConfParserEnabled, |
35 | | SCAppLayerParserRegisterLogger, SCAppLayerProtoDetectConfProtoDetectionEnabled, |
36 | | }; |
37 | | |
38 | | use std; |
39 | | use std::collections::VecDeque; |
40 | | use std::ffi::CString; |
41 | | use std::os::raw::{c_char, c_int, c_void}; |
42 | | |
43 | | pub(super) static mut ALPROTO_WEBSOCKET: AppProto = ALPROTO_UNKNOWN; |
44 | | |
45 | | static mut WEBSOCKET_MAX_PAYLOAD_SIZE: u32 = 0xFFFF; |
46 | | |
47 | | const WEBSOCKET_DECOMPRESS_BUF_SIZE: usize = 8192; |
48 | | |
49 | | #[derive(AppLayerFrameType)] |
50 | | pub enum WebSocketFrameType { |
51 | | Header, |
52 | | Pdu, |
53 | | Data, |
54 | | } |
55 | | |
56 | | #[derive(AppLayerEvent)] |
57 | | pub enum WebSocketEvent { |
58 | | SkipEndOfPayload, |
59 | | ReassemblyLimitReached, |
60 | | } |
61 | | |
62 | | #[derive(Default)] |
63 | | pub struct WebSocketTransaction { |
64 | | tx_id: u64, |
65 | | pub pdu: parser::WebSocketPdu, |
66 | | tx_data: AppLayerTxData, |
67 | | } |
68 | | |
69 | | impl WebSocketTransaction { |
70 | 9.54M | pub fn new(direction: Direction) -> WebSocketTransaction { |
71 | 9.54M | Self { |
72 | 9.54M | tx_data: AppLayerTxData::for_direction(direction), |
73 | 9.54M | ..Default::default() |
74 | 9.54M | } |
75 | 9.54M | } |
76 | | } |
77 | | |
78 | | impl Transaction for WebSocketTransaction { |
79 | 9.46M | fn id(&self) -> u64 { |
80 | 9.46M | self.tx_id |
81 | 9.46M | } |
82 | | } |
83 | | |
84 | | #[derive(Default)] |
85 | | struct WebSocketReassemblyBuffer { |
86 | | data: Vec<u8>, |
87 | | compress: bool, |
88 | | } |
89 | | |
90 | | #[derive(Default)] |
91 | | pub struct WebSocketState { |
92 | | state_data: AppLayerStateData, |
93 | | tx_id: u64, |
94 | | transactions: VecDeque<WebSocketTransaction>, |
95 | | |
96 | | c2s_dec: Option<flate2::Decompress>, |
97 | | s2c_dec: Option<flate2::Decompress>, |
98 | | |
99 | | c2s_buf: WebSocketReassemblyBuffer, |
100 | | s2c_buf: WebSocketReassemblyBuffer, |
101 | | |
102 | | to_skip_tc: u64, |
103 | | to_skip_ts: u64, |
104 | | } |
105 | | |
106 | | impl State<WebSocketTransaction> for WebSocketState { |
107 | 4.74M | fn get_transaction_count(&self) -> usize { |
108 | 4.74M | self.transactions.len() |
109 | 4.74M | } |
110 | | |
111 | 4.73M | fn get_transaction_by_index(&self, index: usize) -> Option<&WebSocketTransaction> { |
112 | 4.73M | self.transactions.get(index) |
113 | 4.73M | } |
114 | | } |
115 | | |
116 | | impl WebSocketState { |
117 | 1.24k | pub fn new() -> Self { |
118 | 1.24k | Default::default() |
119 | 1.24k | } |
120 | | |
121 | | // Free a transaction by ID. |
122 | 4.73M | fn free_tx(&mut self, tx_id: u64) { |
123 | 4.73M | let len = self.transactions.len(); |
124 | 4.73M | let mut found = false; |
125 | 4.73M | let mut index = 0; |
126 | 4.73M | for i in 0..len { |
127 | 4.73M | let tx = &self.transactions[i]; |
128 | 4.73M | if tx.tx_id == tx_id + 1 { |
129 | 4.73M | found = true; |
130 | 4.73M | index = i; |
131 | 4.73M | break; |
132 | 0 | } |
133 | | } |
134 | 4.73M | if found { |
135 | 4.73M | self.transactions.remove(index); |
136 | 4.73M | } |
137 | 4.73M | } |
138 | | |
139 | 0 | pub fn get_tx(&mut self, tx_id: u64) -> Option<&WebSocketTransaction> { |
140 | 0 | self.transactions.iter().find(|tx| tx.tx_id == tx_id + 1) |
141 | 0 | } |
142 | | |
143 | 9.54M | fn new_tx(&mut self, direction: Direction) -> WebSocketTransaction { |
144 | 9.54M | let mut tx = WebSocketTransaction::new(direction); |
145 | 9.54M | self.tx_id += 1; |
146 | 9.54M | tx.tx_id = self.tx_id; |
147 | 9.54M | return tx; |
148 | 9.54M | } |
149 | | |
150 | 56.4k | fn parse( |
151 | 56.4k | &mut self, stream_slice: StreamSlice, direction: Direction, flow: *mut Flow, |
152 | 56.4k | ) -> AppLayerResult { |
153 | 56.4k | let to_skip = if direction == Direction::ToClient { |
154 | 27.7k | &mut self.to_skip_tc |
155 | | } else { |
156 | 28.6k | &mut self.to_skip_ts |
157 | | }; |
158 | 56.4k | let input = stream_slice.as_slice(); |
159 | 56.4k | let mut start = input; |
160 | 56.4k | if *to_skip > 0 { |
161 | 2.24k | if *to_skip >= input.len() as u64 { |
162 | 2.22k | *to_skip -= input.len() as u64; |
163 | 2.22k | return AppLayerResult::ok(); |
164 | 21 | } else { |
165 | 21 | start = &input[*to_skip as usize..]; |
166 | 21 | *to_skip = 0; |
167 | 21 | } |
168 | 54.1k | } |
169 | | |
170 | 54.1k | let max_pl_size = unsafe { WEBSOCKET_MAX_PAYLOAD_SIZE }; |
171 | 9.60M | while !start.is_empty() { |
172 | 9.59M | match parser::parse_message(start, max_pl_size) { |
173 | 9.54M | Ok((rem, pdu)) => { |
174 | 9.54M | let mut tx = self.new_tx(direction); |
175 | 9.54M | let _pdu = Frame::new( |
176 | 9.54M | flow, |
177 | 9.54M | &stream_slice, |
178 | 9.54M | start, |
179 | 9.54M | (start.len() - rem.len() - pdu.payload.len()) as i64, |
180 | 9.54M | WebSocketFrameType::Header as u8, |
181 | 9.54M | Some(tx.tx_id), |
182 | | ); |
183 | 9.54M | let _pdu = Frame::new( |
184 | 9.54M | flow, |
185 | 9.54M | &stream_slice, |
186 | 9.54M | start, |
187 | 9.54M | (start.len() - rem.len()) as i64, |
188 | 9.54M | WebSocketFrameType::Pdu as u8, |
189 | 9.54M | Some(tx.tx_id), |
190 | | ); |
191 | 9.54M | let _pdu = Frame::new( |
192 | 9.54M | flow, |
193 | 9.54M | &stream_slice, |
194 | 9.54M | &start[(start.len() - rem.len() - pdu.payload.len())..], |
195 | 9.54M | pdu.payload.len() as i64, |
196 | 9.54M | WebSocketFrameType::Data as u8, |
197 | 9.54M | Some(tx.tx_id), |
198 | | ); |
199 | 9.54M | start = rem; |
200 | 9.54M | if pdu.to_skip > 0 { |
201 | 512 | if direction == Direction::ToClient { |
202 | 148 | self.to_skip_tc = pdu.to_skip; |
203 | 364 | } else { |
204 | 364 | self.to_skip_ts = pdu.to_skip; |
205 | 364 | } |
206 | 512 | tx.tx_data.set_event(WebSocketEvent::SkipEndOfPayload as u8); |
207 | 9.54M | } |
208 | 9.54M | if pdu.compress { |
209 | | // RFC 7692 section 7.1.2 states that |
210 | | // absence of precision means LZ77 sliding window of up to 2^15 bytes |
211 | 33.7k | if direction == Direction::ToClient && self.s2c_dec.is_none() { |
212 | 412 | self.s2c_dec = Some(Decompress::new_with_window_bits(false, 15)); |
213 | 33.2k | } else if direction == Direction::ToServer && self.c2s_dec.is_none() { |
214 | 797 | self.c2s_dec = Some(Decompress::new_with_window_bits(false, 15)); |
215 | 32.4k | } |
216 | 9.51M | } |
217 | 9.54M | let (buf, dec) = if direction == Direction::ToClient { |
218 | 4.53M | (&mut self.s2c_buf, &mut self.s2c_dec) |
219 | | } else { |
220 | 5.01M | (&mut self.c2s_buf, &mut self.c2s_dec) |
221 | | }; |
222 | 9.54M | let mut compress = pdu.compress; |
223 | 9.54M | if pdu.opcode < 8 && (!buf.data.is_empty() || !pdu.fin) { |
224 | 9.42M | if buf.data.is_empty() { |
225 | 2.69M | buf.compress = pdu.compress; |
226 | 6.73M | } |
227 | 9.42M | if buf.data.len() + pdu.payload.len() < max_pl_size as usize { |
228 | 8.65M | buf.data.extend(&pdu.payload); |
229 | 8.65M | } else if buf.data.len() < max_pl_size as usize { |
230 | 280 | buf.data |
231 | 280 | .extend(&pdu.payload[..max_pl_size as usize - buf.data.len()]); |
232 | 280 | tx.tx_data |
233 | 280 | .set_event(WebSocketEvent::ReassemblyLimitReached as u8); |
234 | 767k | } |
235 | 125k | } |
236 | 9.54M | tx.pdu = pdu; |
237 | 9.54M | if tx.pdu.opcode < 8 && tx.pdu.fin && !buf.data.is_empty() { |
238 | 29.4k | // the final PDU gets the full reassembled payload |
239 | 29.4k | compress = buf.compress; |
240 | 29.4k | std::mem::swap(&mut tx.pdu.payload, &mut buf.data); |
241 | 29.4k | buf.data.clear(); |
242 | 9.51M | } |
243 | 9.54M | if compress && tx.pdu.fin { |
244 | 17.1k | buf.compress = false; |
245 | | // cf RFC 7692 section-7.2.2 |
246 | 17.1k | tx.pdu.payload.extend_from_slice(&[0, 0, 0xFF, 0xFF]); |
247 | 17.1k | let mut v = Vec::with_capacity(std::cmp::min( |
248 | | WEBSOCKET_DECOMPRESS_BUF_SIZE, |
249 | | // Do not allocate 8kbytes for a small size. |
250 | | // Numbers here may be optimized. |
251 | 17.1k | 256 + 16 * tx.pdu.payload.len(), |
252 | | )); |
253 | 17.1k | if let Some(dec) = dec { |
254 | 17.1k | let expect = dec.total_in() + tx.pdu.payload.len() as u64; |
255 | 17.1k | let start = dec.total_in(); |
256 | 17.1k | let mut e = dec.decompress_vec( |
257 | 17.1k | &tx.pdu.payload, |
258 | 17.1k | &mut v, |
259 | 17.1k | FlushDecompress::Finish, |
260 | | ); |
261 | 17.6k | while e.is_ok() && dec.total_in() < expect { |
262 | 1.91k | let mut s = vec![0u8; WEBSOCKET_DECOMPRESS_BUF_SIZE]; |
263 | 1.91k | let before = dec.total_out(); |
264 | 1.91k | let check = dec.total_in(); |
265 | 1.91k | e = dec.decompress( |
266 | 1.91k | &tx.pdu.payload[(dec.total_in() - start) as usize..], |
267 | 1.91k | &mut s, |
268 | 1.91k | FlushDecompress::Finish, |
269 | 1.91k | ); |
270 | 1.91k | if v.len() < max_pl_size as usize { |
271 | 1.62k | let end = if v.len() + (dec.total_out() - before) as usize |
272 | 1.62k | > max_pl_size as usize |
273 | | { |
274 | 26 | max_pl_size as usize - v.len() |
275 | | } else { |
276 | 1.60k | (dec.total_out() - before) as usize |
277 | | }; |
278 | 1.62k | v.extend_from_slice(&s[..end]); |
279 | 283 | } |
280 | 1.91k | if check >= dec.total_in() { |
281 | | // safety check against infinite loop : dec.total_in() should increase |
282 | 1.34k | break; |
283 | 571 | } |
284 | | } |
285 | 17.1k | if !v.is_empty() { |
286 | 5.72k | std::mem::swap(&mut tx.pdu.payload, &mut v); |
287 | 11.3k | } |
288 | 0 | } |
289 | 9.53M | } |
290 | 9.54M | if tx.pdu.fin { |
291 | 145k | sc_app_layer_parser_trigger_raw_stream_inspection(flow, direction as i32); |
292 | 9.40M | } |
293 | 9.54M | self.transactions.push_back(tx); |
294 | | } |
295 | 48.7k | Err(nom::Err::Incomplete(needed)) => { |
296 | 48.7k | if let Needed::Size(n) = needed { |
297 | 48.7k | let n = usize::from(n); |
298 | | // Not enough data. just ask for one more byte. |
299 | 48.7k | let consumed = input.len() - start.len(); |
300 | 48.7k | let needed = start.len() + n; |
301 | 48.7k | return AppLayerResult::incomplete(consumed as u32, needed as u32); |
302 | 0 | } |
303 | 0 | return AppLayerResult::err(); |
304 | | } |
305 | | Err(_) => { |
306 | 0 | return AppLayerResult::err(); |
307 | | } |
308 | | } |
309 | | } |
310 | | // Input was fully consumed. |
311 | 5.40k | return AppLayerResult::ok(); |
312 | 56.4k | } |
313 | | } |
314 | | |
315 | | // C exports. |
316 | | |
317 | 0 | unsafe extern "C" fn websocket_probing_parser( |
318 | 0 | _flow: *const Flow, _direction: u8, input: *const u8, input_len: u32, _rdir: *mut u8, |
319 | 0 | ) -> AppProto { |
320 | 0 | if !input.is_null() { |
321 | 0 | let slice = build_slice!(input, input_len as usize); |
322 | 0 | if !slice.is_empty() { |
323 | | // just check reserved bits are zeroed, except RSV1 |
324 | | // as RSV1 is used for compression cf RFC 7692 |
325 | 0 | if slice[0] & 0x30 == 0 { |
326 | 0 | return ALPROTO_WEBSOCKET; |
327 | 0 | } |
328 | 0 | return ALPROTO_FAILED; |
329 | 0 | } |
330 | 0 | } |
331 | 0 | return ALPROTO_UNKNOWN; |
332 | 0 | } |
333 | | |
334 | 1.24k | extern "C" fn websocket_state_new(_orig_state: *mut c_void, _orig_proto: AppProto) -> *mut c_void { |
335 | 1.24k | let state = WebSocketState::new(); |
336 | 1.24k | let boxed = Box::new(state); |
337 | 1.24k | return Box::into_raw(boxed) as *mut c_void; |
338 | 1.24k | } |
339 | | |
340 | 1.24k | unsafe extern "C" fn websocket_state_free(state: *mut c_void) { |
341 | 1.24k | std::mem::drop(Box::from_raw(state as *mut WebSocketState)); |
342 | 1.24k | } |
343 | | |
344 | 4.73M | unsafe extern "C" fn websocket_state_tx_free(state: *mut c_void, tx_id: u64) { |
345 | 4.73M | let state = cast_pointer!(state, WebSocketState); |
346 | 4.73M | state.free_tx(tx_id); |
347 | 4.73M | } |
348 | | |
349 | 28.6k | unsafe extern "C" fn websocket_parse_request( |
350 | 28.6k | flow: *mut Flow, state: *mut c_void, _pstate: *mut AppLayerParserState, |
351 | 28.6k | stream_slice: StreamSlice, _data: *const c_void, |
352 | 28.6k | ) -> AppLayerResult { |
353 | 28.6k | let state = cast_pointer!(state, WebSocketState); |
354 | 28.6k | state.parse(stream_slice, Direction::ToServer, flow) |
355 | 28.6k | } |
356 | | |
357 | 27.7k | unsafe extern "C" fn websocket_parse_response( |
358 | 27.7k | flow: *mut Flow, state: *mut c_void, _pstate: *mut AppLayerParserState, |
359 | 27.7k | stream_slice: StreamSlice, _data: *const c_void, |
360 | 27.7k | ) -> AppLayerResult { |
361 | 27.7k | let state = cast_pointer!(state, WebSocketState); |
362 | 27.7k | state.parse(stream_slice, Direction::ToClient, flow) |
363 | 27.7k | } |
364 | | |
365 | 0 | unsafe extern "C" fn websocket_state_get_tx(state: *mut c_void, tx_id: u64) -> *mut c_void { |
366 | 0 | let state = cast_pointer!(state, WebSocketState); |
367 | 0 | match state.get_tx(tx_id) { |
368 | 0 | Some(tx) => { |
369 | 0 | return tx as *const _ as *mut _; |
370 | | } |
371 | | None => { |
372 | 0 | return std::ptr::null_mut(); |
373 | | } |
374 | | } |
375 | 0 | } |
376 | | |
377 | 168k | unsafe extern "C" fn websocket_state_get_tx_count(state: *mut c_void) -> u64 { |
378 | 168k | let state = cast_pointer!(state, WebSocketState); |
379 | 168k | return state.tx_id; |
380 | 168k | } |
381 | | |
382 | 9.46M | unsafe extern "C" fn websocket_tx_get_alstate_progress(_tx: *mut c_void, _direction: u8) -> c_int { |
383 | 9.46M | return 1; |
384 | 9.46M | } |
385 | | |
386 | | export_tx_data_get!(websocket_get_tx_data, WebSocketTransaction); |
387 | | export_state_data_get!(websocket_get_state_data, WebSocketState); |
388 | | |
389 | | // Parser name as a C style string. |
390 | | const PARSER_NAME: &[u8] = b"websocket\0"; |
391 | | |
392 | | #[no_mangle] |
393 | 40 | pub unsafe extern "C" fn SCRegisterWebSocketParser() { |
394 | 40 | let parser = RustParser { |
395 | 40 | name: PARSER_NAME.as_ptr() as *const c_char, |
396 | 40 | default_port: std::ptr::null(), |
397 | 40 | ipproto: IPPROTO_TCP, |
398 | 40 | probe_ts: Some(websocket_probing_parser), |
399 | 40 | probe_tc: Some(websocket_probing_parser), |
400 | 40 | min_depth: 0, |
401 | 40 | max_depth: 16, |
402 | 40 | state_new: websocket_state_new, |
403 | 40 | state_free: websocket_state_free, |
404 | 40 | tx_free: websocket_state_tx_free, |
405 | 40 | parse_ts: websocket_parse_request, |
406 | 40 | parse_tc: websocket_parse_response, |
407 | 40 | get_tx_count: websocket_state_get_tx_count, |
408 | 40 | get_tx: websocket_state_get_tx, |
409 | 40 | tx_comp_st_ts: 1, |
410 | 40 | tx_comp_st_tc: 1, |
411 | 40 | tx_get_progress: websocket_tx_get_alstate_progress, |
412 | 40 | get_eventinfo: Some(WebSocketEvent::get_event_info), |
413 | 40 | get_eventinfo_byid: Some(WebSocketEvent::get_event_info_by_id), |
414 | 40 | localstorage_new: None, |
415 | 40 | localstorage_free: None, |
416 | 40 | get_tx_files: None, |
417 | 40 | get_tx_iterator: Some( |
418 | 40 | applayer::state_get_tx_iterator::<WebSocketState, WebSocketTransaction>, |
419 | 40 | ), |
420 | 40 | get_tx_data: websocket_get_tx_data, |
421 | 40 | get_state_data: websocket_get_state_data, |
422 | 40 | apply_tx_config: None, |
423 | 40 | flags: 0, // do not accept gaps as there is no good way to resync |
424 | 40 | get_frame_id_by_name: Some(WebSocketFrameType::ffi_id_from_name), |
425 | 40 | get_frame_name_by_id: Some(WebSocketFrameType::ffi_name_from_id), |
426 | 40 | get_state_id_by_name: None, |
427 | 40 | get_state_name_by_id: None, |
428 | 40 | }; |
429 | | |
430 | 40 | let ip_proto_str = CString::new("tcp").unwrap(); |
431 | | |
432 | 40 | if SCAppLayerProtoDetectConfProtoDetectionEnabled(ip_proto_str.as_ptr(), parser.name) != 0 { |
433 | 40 | let alproto = AppLayerRegisterProtocolDetection(&parser, 1); |
434 | 40 | ALPROTO_WEBSOCKET = alproto; |
435 | 40 | if SCAppLayerParserConfParserEnabled(ip_proto_str.as_ptr(), parser.name) != 0 { |
436 | 40 | let _ = AppLayerRegisterParser(&parser, alproto); |
437 | 40 | } |
438 | | SCLogDebug!("Rust websocket parser registered."); |
439 | 40 | if let Some(val) = conf_get("app-layer.protocols.websocket.max-payload-size") { |
440 | 0 | if let Ok(v) = val.parse::<u32>() { |
441 | 0 | WEBSOCKET_MAX_PAYLOAD_SIZE = v; |
442 | 0 | } else { |
443 | 0 | SCLogError!("Invalid value for websocket.max-payload-size"); |
444 | | } |
445 | 40 | } |
446 | 40 | SCAppLayerParserRegisterLogger(IPPROTO_TCP, ALPROTO_WEBSOCKET); |
447 | 0 | } else { |
448 | 0 | SCLogDebug!("Protocol detector and parser disabled for WEBSOCKET."); |
449 | 0 | } |
450 | 40 | } |