Coverage Report

Created: 2026-09-06 07:25

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
/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
}