Skip to main content

retina_core/protocols/stream/tls/
parser.rs

1//! TLS handshake parser.
2//!
3//! The TLS handshake parser uses a [fork](https://github.com/thegwan/tls-parser) of the
4//! [tls-parser](https://docs.rs/tls-parser/latest/tls_parser/) crate to parse the handshake phase
5//! of a TLS connection. It maintains TLS state, stores selected parameters, and handles
6//! defragmentation.
7//!
8//! Adapted from [the Rusticata TLS
9//! parser](https://github.com/rusticata/rusticata/blob/master/src/tls.rs).
10
11use super::handshake::{
12    Certificate, ClientDHParams, ClientECDHParams, ClientHello, ClientKeyExchange, ClientRSAParams,
13    KeyShareEntry, ServerDHParams, ServerECDHParams, ServerHello, ServerKeyExchange,
14    ServerRSAParams,
15};
16use super::Tls;
17use crate::conntrack::conn::conn_info::ConnState;
18use crate::conntrack::pdu::L4Pdu;
19use crate::protocols::stream::{ConnParsable, ParseResult, ProbeResult, Session, SessionData};
20
21use tls_parser::*;
22
23/// Parses a single TLS handshake per connection.
24#[derive(Debug)]
25pub struct TlsParser {
26    sessions: Vec<Tls>,
27}
28
29impl TlsParser {}
30
31impl Default for TlsParser {
32    fn default() -> Self {
33        TlsParser {
34            sessions: vec![Tls::new()],
35        }
36    }
37}
38
39impl ConnParsable for TlsParser {
40    fn parse(&mut self, pdu: &L4Pdu) -> ParseResult {
41        log::debug!("Updating parser tls");
42        let offset = pdu.offset();
43        let length = pdu.length();
44        if length == 0 {
45            return ParseResult::Skipped;
46        }
47
48        if let Ok(data) = (pdu.mbuf_ref()).get_data_slice(offset, length) {
49            self.sessions[0].parse_tcp_level(data, pdu.dir)
50        } else {
51            log::warn!("Malformed packet");
52            ParseResult::Skipped
53        }
54    }
55
56    fn probe(&self, pdu: &L4Pdu) -> ProbeResult {
57        if pdu.length() <= 2 {
58            return ProbeResult::Unsure;
59        }
60
61        let offset = pdu.offset();
62        let length = pdu.length();
63        if let Ok(data) = (pdu.mbuf_ref()).get_data_slice(offset, length) {
64            // First byte is record type (between 0x14 and 0x17, 0x16 is handhake) Second is TLS
65            // version major (0x3) Third is TLS version minor (0x0 for SSLv3, 0x1 for TLSv1.0, etc.)
66            // Does not support versions <= SSLv2
67            match (data[0], data[1], data[2]) {
68                (0x14..=0x17, 0x03, 0..=3) => ProbeResult::Certain,
69                _ => ProbeResult::NotForUs,
70            }
71        } else {
72            log::warn!("Malformed packet");
73            ProbeResult::Error
74        }
75    }
76
77    fn remove_session(&mut self, _session_id: usize) -> Option<Session> {
78        self.sessions.pop().map(|tls| Session {
79            data: SessionData::Tls(Box::new(tls)),
80            id: 0,
81        })
82    }
83
84    fn drain_sessions(&mut self) -> Vec<Session> {
85        self.sessions
86            .drain(..)
87            .map(|tls| Session {
88                data: SessionData::Tls(Box::new(tls)),
89                id: 0,
90            })
91            .collect()
92    }
93
94    fn session_match_state(&self) -> ConnState {
95        ConnState::Remove
96    }
97
98    fn session_nomatch_state(&self) -> ConnState {
99        ConnState::Remove
100    }
101}
102
103// ------------------------------------------------------------
104
105impl Tls {
106    /// Allocate a new TLS handshake instance.
107    pub(crate) fn new() -> Tls {
108        Tls {
109            client_hello: None,
110            server_hello: None,
111            server_certificates: vec![],
112            client_certificates: vec![],
113            server_key_exchange: None,
114            client_key_exchange: None,
115            state: TlsState::None,
116            tcp_buffer: vec![],
117            record_buffer: vec![],
118        }
119    }
120
121    /// Parse a ClientHello message.
122    pub(crate) fn parse_handshake_clienthello(&mut self, content: &TlsClientHelloContents) {
123        let mut client_hello = ClientHello {
124            version: content.version,
125            random: content.random.to_vec(),
126            session_id: match content.session_id {
127                Some(v) => v.to_vec(),
128                None => vec![],
129            },
130            cipher_suites: content.ciphers.to_vec(),
131            compression_algs: content.comp.to_vec(),
132            ..ClientHello::default()
133        };
134
135        let ext = parse_tls_client_hello_extensions(content.ext.unwrap_or(b""));
136        log::trace!("client extensions: {:#?}", ext);
137        match &ext {
138            Ok((rem, ref ext_lst)) => {
139                if !rem.is_empty() {
140                    log::debug!("warn: extensions not entirely parsed");
141                }
142                for extension in ext_lst {
143                    client_hello
144                        .extension_list
145                        .push(TlsExtensionType::from(extension));
146                    match *extension {
147                        TlsExtension::SNI(ref v) if !v.is_empty() => {
148                            let sni = v[0].1;
149                            client_hello.server_name = Some(match std::str::from_utf8(sni) {
150                                Ok(name) => name.to_string(),
151                                Err(_) => format!("<Invalid UTF-8: {}>", hex::encode(sni)),
152                            });
153                        }
154                        TlsExtension::SupportedGroups(ref v) => {
155                            client_hello.supported_groups = v.clone();
156                        }
157                        TlsExtension::EcPointFormats(v) => {
158                            client_hello.ec_point_formats = v.to_vec();
159                        }
160                        TlsExtension::SignatureAlgorithms(ref v) => {
161                            client_hello.signature_algs = v.clone();
162                        }
163                        TlsExtension::ALPN(ref v) => {
164                            for proto in v {
165                                client_hello.alpn_protocols.push(
166                                    match std::str::from_utf8(proto) {
167                                        Ok(proto) => proto.to_string(),
168                                        Err(_) => {
169                                            format!("<Invalid UTF-8: {}>", hex::encode(proto))
170                                        }
171                                    },
172                                );
173                            }
174                        }
175                        TlsExtension::KeyShare(ref v) => {
176                            log::debug!("Client Shares: {:?}", v);
177                            client_hello.key_shares = v
178                                .iter()
179                                .map(|k| KeyShareEntry {
180                                    group: k.group,
181                                    kx_data: k.kx.to_vec(),
182                                })
183                                .collect();
184                        }
185                        TlsExtension::SupportedVersions(ref v) => {
186                            client_hello.supported_versions = v.clone();
187                        }
188                        _ => (),
189                    }
190                }
191            }
192            e => log::debug!("Could not parse extensions: {:?}", e),
193        };
194        self.client_hello = Some(client_hello);
195    }
196
197    /// Parse a ServerHello message.
198    fn parse_handshake_serverhello(&mut self, content: &TlsServerHelloContents) {
199        let mut server_hello = ServerHello {
200            version: content.version,
201            random: content.random.to_vec(),
202            session_id: match content.session_id {
203                Some(v) => v.to_vec(),
204                None => vec![],
205            },
206            cipher_suite: content.cipher,
207            compression_alg: content.compression,
208            ..ServerHello::default()
209        };
210
211        let ext = parse_tls_server_hello_extensions(content.ext.unwrap_or(b""));
212        log::debug!("server_hello extensions: {:#?}", ext);
213        match &ext {
214            Ok((rem, ref ext_lst)) => {
215                if !rem.is_empty() {
216                    log::debug!("warn: extensions not entirely parsed");
217                }
218                for extension in ext_lst {
219                    server_hello
220                        .extension_list
221                        .push(TlsExtensionType::from(extension));
222                    match *extension {
223                        TlsExtension::EcPointFormats(v) => {
224                            server_hello.ec_point_formats = v.to_vec();
225                        }
226                        TlsExtension::ALPN(ref v) if !v.is_empty() => {
227                            server_hello.alpn_protocol = Some(match std::str::from_utf8(v[0]) {
228                                Ok(proto) => proto.to_string(),
229                                Err(_) => format!("<Invalid UTF-8: {}>", hex::encode(v[0])),
230                            });
231                        }
232                        TlsExtension::KeyShare(ref v) => {
233                            log::debug!("Server Share: {:?}", v);
234                            if !v.is_empty() {
235                                server_hello.key_share = Some(KeyShareEntry {
236                                    group: v[0].group,
237                                    kx_data: v[0].kx.to_vec(),
238                                });
239                            }
240                        }
241                        TlsExtension::SupportedVersions(ref v) if !v.is_empty() => {
242                            server_hello.selected_version = Some(v[0]);
243                        }
244                        _ => (),
245                    }
246                }
247            }
248            e => log::debug!("Could not parse extensions: {:?}", e),
249        };
250        self.server_hello = Some(server_hello);
251    }
252
253    /// Parse a Certificate message.
254    fn parse_handshake_certificate(&mut self, content: &TlsCertificateContents, direction: bool) {
255        log::trace!("cert chain length: {}", content.cert_chain.len());
256        if direction {
257            // client -> server
258            for cert in &content.cert_chain {
259                self.client_certificates.push(Certificate {
260                    raw: cert.data.to_vec(),
261                })
262            }
263        } else {
264            // server -> client
265            for cert in &content.cert_chain {
266                self.server_certificates.push(Certificate {
267                    raw: cert.data.to_vec(),
268                })
269            }
270        }
271    }
272
273    /// Parse a ServerKeyExchange message.
274    fn parse_handshake_serverkeyexchange(&mut self, content: &TlsServerKeyExchangeContents) {
275        log::trace!("SKE: {:?}", content);
276        if let Some(cipher) = self.cipher_suite() {
277            match &cipher.kx {
278                TlsCipherKx::Ecdhe | TlsCipherKx::Ecdh => {
279                    if let Ok((_sig, ref parsed)) = parse_server_ecdh_params(content.parameters) {
280                        if let ECParametersContent::NamedGroup(curve) =
281                            parsed.curve_params.params_content
282                        {
283                            let ecdh_params = ServerECDHParams {
284                                curve,
285                                kx_data: parsed.public.point.to_vec(),
286                            };
287                            self.server_key_exchange = Some(ServerKeyExchange::Ecdh(ecdh_params));
288                        };
289                    }
290                }
291                TlsCipherKx::Dhe | TlsCipherKx::Dh => {
292                    if let Ok((_sig, ref parsed)) = parse_server_dh_params(content.parameters) {
293                        let dh_params = ServerDHParams {
294                            prime: parsed.dh_p.to_vec(),
295                            generator: parsed.dh_g.to_vec(),
296                            kx_data: parsed.dh_ys.to_vec(),
297                        };
298                        self.server_key_exchange = Some(ServerKeyExchange::Dh(dh_params));
299                    }
300                }
301                TlsCipherKx::Rsa => {
302                    if let Ok((_sig, ref parsed)) = parse_server_rsa_params(content.parameters) {
303                        let rsa_params = ServerRSAParams {
304                            modulus: parsed.modulus.to_vec(),
305                            exponent: parsed.exponent.to_vec(),
306                        };
307                        self.server_key_exchange = Some(ServerKeyExchange::Rsa(rsa_params));
308                    }
309                }
310                _ => {
311                    self.server_key_exchange =
312                        Some(ServerKeyExchange::Unknown(content.parameters.to_vec()))
313                }
314            }
315        }
316    }
317
318    /// Parse a ClientKeyExchange message.
319    fn parse_handshake_clientkeyexchange(&mut self, content: &TlsClientKeyExchangeContents) {
320        log::trace!("CKE: {:?}", content);
321        if let Some(cipher) = self.cipher_suite() {
322            match &cipher.kx {
323                TlsCipherKx::Ecdhe | TlsCipherKx::Ecdh => {
324                    if let Ok((_rem, ref parsed)) = parse_client_ecdh_params(content.parameters) {
325                        let ecdh_params = ClientECDHParams {
326                            kx_data: parsed.ecdh_yc.point.to_vec(),
327                        };
328                        self.client_key_exchange = Some(ClientKeyExchange::Ecdh(ecdh_params));
329                    }
330                }
331                TlsCipherKx::Dhe | TlsCipherKx::Dh => {
332                    if let Ok((_rem, ref parsed)) = parse_client_dh_params(content.parameters) {
333                        let dh_params = ClientDHParams {
334                            kx_data: parsed.dh_yc.to_vec(),
335                        };
336                        self.client_key_exchange = Some(ClientKeyExchange::Dh(dh_params));
337                    }
338                }
339                TlsCipherKx::Rsa => {
340                    if let Ok((_rem, ref parsed)) = parse_client_rsa_params(content.parameters) {
341                        let rsa_params = ClientRSAParams {
342                            encrypted_pms: parsed.data.to_vec(),
343                        };
344                        self.client_key_exchange = Some(ClientKeyExchange::Rsa(rsa_params));
345                    }
346                }
347                _ => {
348                    self.client_key_exchange =
349                        Some(ClientKeyExchange::Unknown(content.parameters.to_vec()))
350                }
351            }
352        }
353        //self.client_key_exchange = Some(client_key_exchange);
354    }
355
356    /// Parse a TLS message.
357    pub(crate) fn parse_message_level(&mut self, msg: &TlsMessage, direction: bool) -> ParseResult {
358        log::trace!("parse_message_level {:?}", msg);
359
360        // do not parse if session is encrypted
361        if self.state == TlsState::ClientChangeCipherSpec {
362            log::trace!("TLS session encrypted, activating bypass");
363            return ParseResult::Done(0);
364        }
365
366        // update state machine
367        match tls_state_transition(self.state, msg, direction) {
368            Ok(s) => self.state = s,
369            Err(_) => {
370                self.state = TlsState::Invalid;
371            }
372        };
373        log::trace!("TLS new state: {:?}", self.state);
374
375        // extract variables
376        match *msg {
377            TlsMessage::Handshake(ref m) => match *m {
378                TlsMessageHandshake::ClientHello(ref content) => {
379                    self.parse_handshake_clienthello(content);
380                }
381                TlsMessageHandshake::ServerHello(ref content) => {
382                    self.parse_handshake_serverhello(content);
383                }
384                TlsMessageHandshake::Certificate(ref content) => {
385                    self.parse_handshake_certificate(content, direction);
386                }
387                TlsMessageHandshake::ServerKeyExchange(ref content) => {
388                    self.parse_handshake_serverkeyexchange(content);
389                }
390                TlsMessageHandshake::ClientKeyExchange(ref content) => {
391                    self.parse_handshake_clientkeyexchange(content);
392                }
393
394                _ => (),
395            },
396            TlsMessage::Alert(ref a) if a.severity == TlsAlertSeverity::Fatal => {
397                return ParseResult::Done(0);
398            }
399            TlsMessage::Alert(_) => {}
400            _ => (),
401        }
402
403        ParseResult::Continue(0)
404    }
405
406    /// Parse a TLS record.
407    pub(crate) fn parse_record_level(
408        &mut self,
409        record: &TlsRawRecord<'_>,
410        direction: bool,
411    ) -> ParseResult {
412        let mut v: Vec<u8>;
413        let mut status = ParseResult::Continue(0);
414
415        log::trace!("parse_record_level ({} bytes)", record.data.len());
416        log::trace!("{:?}", record.hdr);
417        // log::trace!("{:?}", record.data);
418
419        // do not parse if session is encrypted
420        if self.state == TlsState::ClientChangeCipherSpec {
421            log::trace!("TLS session encrypted, activating bypass");
422            return ParseResult::Done(0);
423        }
424
425        // only parse some message types (the Content type, first byte of TLS record)
426        match record.hdr.record_type {
427            TlsRecordType::ChangeCipherSpec => (),
428            TlsRecordType::Handshake => (),
429            TlsRecordType::Alert => (),
430            _ => return ParseResult::Continue(0),
431        }
432
433        // Check if a record is being defragmented
434        let record_buffer = match self.record_buffer.len() {
435            0 => record.data,
436            _ => {
437                // sanity check vector length to avoid memory exhaustion maximum length may be 2^24
438                // (handshake message)
439                if self.record_buffer.len() + record.data.len() > 16_777_216 {
440                    return ParseResult::Skipped;
441                };
442                v = self.record_buffer.split_off(0);
443                v.extend_from_slice(record.data);
444                v.as_slice()
445            }
446        };
447
448        // TODO: record may be compressed Parse record contents as plaintext
449        match parse_tls_record_with_header(record_buffer, &record.hdr) {
450            Ok((rem, ref msg_list)) => {
451                for msg in msg_list {
452                    status = self.parse_message_level(msg, direction);
453                    if status != ParseResult::Continue(0) {
454                        return status;
455                    }
456                }
457                if !rem.is_empty() {
458                    log::debug!("warn: extra bytes in TLS record: {:?}", rem);
459                };
460            }
461            Err(Err::Incomplete(needed)) => {
462                log::trace!(
463                    "Defragmentation required (TLS record), missing {:?} bytes",
464                    needed
465                );
466                self.record_buffer.extend_from_slice(record.data);
467            }
468            Err(_e) => {
469                log::debug!("warn: parse_tls_record_with_header failed");
470                return ParseResult::Skipped;
471            }
472        };
473
474        status
475    }
476
477    /// Parse a TCP segment, handling TCP chunks fragmentation.
478    pub(crate) fn parse_tcp_level(&mut self, data: &[u8], direction: bool) -> ParseResult {
479        let mut v: Vec<u8>;
480        let mut status = ParseResult::Continue(0);
481        log::trace!("parse_tcp_level ({} bytes)", data.len());
482        log::trace!("defrag buffer size: {}", self.tcp_buffer.len());
483
484        // do not parse if session is encrypted
485        if self.state == TlsState::ClientChangeCipherSpec {
486            log::trace!("TLS session encrypted, activating bypass");
487            return ParseResult::Done(0);
488        };
489        // Check if TCP data is being defragmented
490        let tcp_buffer = match self.tcp_buffer.len() {
491            0 => data,
492            _ => {
493                // sanity check vector length to avoid memory exhaustion maximum length may be 2^24
494                // (handshake message)
495                if self.tcp_buffer.len() + data.len() > 16_777_216 {
496                    return ParseResult::Skipped;
497                };
498                v = self.tcp_buffer.split_off(0);
499                v.extend_from_slice(data);
500                v.as_slice()
501            }
502        };
503        let mut cur_data = tcp_buffer;
504        while !cur_data.is_empty() {
505            // parse each TLS record in the TCP segment (there could be multiple)
506            match parse_tls_raw_record(cur_data) {
507                Ok((rem, ref record)) => {
508                    cur_data = rem;
509                    status = self.parse_record_level(record, direction);
510                    if status != ParseResult::Continue(0) {
511                        return status;
512                    }
513                }
514                Err(Err::Incomplete(needed)) => {
515                    log::trace!(
516                        "Defragmentation required (TCP level), missing {:?} bytes",
517                        needed
518                    );
519                    self.tcp_buffer.extend_from_slice(cur_data);
520                    break;
521                }
522                Err(_e) => {
523                    log::debug!("warn: Parsing raw record failed");
524                    break;
525                }
526            }
527        }
528        status
529    }
530}