Skip to main content

rustls/msgs/
persist.rs

1use alloc::vec::Vec;
2use core::cmp;
3
4use pki_types::{DnsName, UnixTime};
5use zeroize::Zeroizing;
6
7use crate::client::ResolvesClientCert;
8use crate::enums::{CipherSuite, ProtocolVersion};
9use crate::error::InvalidMessage;
10use crate::msgs::base::{MaybeEmpty, PayloadU8, PayloadU16};
11use crate::msgs::codec::{Codec, Reader};
12#[cfg(feature = "tls12")]
13use crate::msgs::handshake::SessionId;
14use crate::msgs::handshake::{CertificateChain, ProtocolName};
15use crate::sync::{Arc, Weak};
16#[cfg(feature = "tls12")]
17use crate::tls12::Tls12CipherSuite;
18use crate::tls13::Tls13CipherSuite;
19use crate::verify::ServerCertVerifier;
20
21pub(crate) struct Retrieved<T> {
22    pub(crate) value: T,
23    retrieved_at: UnixTime,
24}
25
26impl<T> Retrieved<T> {
27    pub(crate) fn new(value: T, retrieved_at: UnixTime) -> Self {
28        Self {
29            value,
30            retrieved_at,
31        }
32    }
33
34    pub(crate) fn map<M>(&self, f: impl FnOnce(&T) -> Option<&M>) -> Option<Retrieved<&M>> {
35        Some(Retrieved {
36            value: f(&self.value)?,
37            retrieved_at: self.retrieved_at,
38        })
39    }
40}
41
42impl Retrieved<&Tls13ClientSessionValue> {
43    pub(crate) fn obfuscated_ticket_age(&self) -> u32 {
44        let age_secs = self
45            .retrieved_at
46            .as_secs()
47            .saturating_sub(self.value.common.epoch);
48        // nb. tickets have an upper age limit of ~7 days, well short of the 49 days here
49        let age_millis = u32::try_from(age_secs)
50            .unwrap_or(u32::MAX)
51            .saturating_mul(1000);
52        age_millis.wrapping_add(self.value.age_add)
53    }
54}
55
56impl<T: core::ops::Deref<Target = ClientSessionCommon>> Retrieved<T> {
57    pub(crate) fn has_expired(&self) -> bool {
58        let common = &*self.value;
59        common.lifetime_secs != 0
60            && common
61                .epoch
62                .saturating_add(u64::from(common.lifetime_secs))
63                < self.retrieved_at.as_secs()
64    }
65}
66
67impl<T> core::ops::Deref for Retrieved<T> {
68    type Target = T;
69
70    fn deref(&self) -> &Self::Target {
71        &self.value
72    }
73}
74
75#[derive(Debug)]
76pub struct Tls13ClientSessionValue {
77    suite: &'static Tls13CipherSuite,
78    age_add: u32,
79    max_early_data_size: u32,
80    pub(crate) common: ClientSessionCommon,
81    quic_params: PayloadU16,
82}
83
84impl Tls13ClientSessionValue {
85    pub(crate) fn new(
86        suite: &'static Tls13CipherSuite,
87        ticket: Arc<PayloadU16>,
88        secret: &[u8],
89        server_cert_chain: CertificateChain<'static>,
90        server_cert_verifier: &Arc<dyn ServerCertVerifier>,
91        client_creds: &Arc<dyn ResolvesClientCert>,
92        time_now: UnixTime,
93        lifetime_secs: u32,
94        age_add: u32,
95        max_early_data_size: u32,
96    ) -> Self {
97        Self {
98            suite,
99            age_add,
100            max_early_data_size,
101            common: ClientSessionCommon::new(
102                ticket,
103                secret,
104                time_now,
105                lifetime_secs,
106                server_cert_chain,
107                server_cert_verifier,
108                client_creds,
109            ),
110            quic_params: PayloadU16::new(Vec::new()),
111        }
112    }
113
114    pub fn max_early_data_size(&self) -> u32 {
115        self.max_early_data_size
116    }
117
118    pub fn suite(&self) -> &'static Tls13CipherSuite {
119        self.suite
120    }
121
122    #[doc(hidden)]
123    /// Test only: rewind epoch by `delta` seconds.
124    pub fn rewind_epoch(&mut self, delta: u32) {
125        self.common.epoch -= delta as u64;
126    }
127
128    #[doc(hidden)]
129    /// Test only: replace `max_early_data_size` with `new`
130    pub fn _private_set_max_early_data_size(&mut self, new: u32) {
131        self.max_early_data_size = new;
132    }
133
134    pub fn set_quic_params(&mut self, quic_params: &[u8]) {
135        self.quic_params = PayloadU16::new(quic_params.to_vec());
136    }
137
138    pub fn quic_params(&self) -> Vec<u8> {
139        self.quic_params.0.clone()
140    }
141}
142
143impl core::ops::Deref for Tls13ClientSessionValue {
144    type Target = ClientSessionCommon;
145
146    fn deref(&self) -> &Self::Target {
147        &self.common
148    }
149}
150
151#[derive(Debug, Clone)]
152pub struct Tls12ClientSessionValue {
153    #[cfg(feature = "tls12")]
154    suite: &'static Tls12CipherSuite,
155    #[cfg(feature = "tls12")]
156    pub(crate) session_id: SessionId,
157    #[cfg(feature = "tls12")]
158    extended_ms: bool,
159    #[doc(hidden)]
160    #[cfg(feature = "tls12")]
161    pub(crate) common: ClientSessionCommon,
162}
163
164#[cfg(feature = "tls12")]
165impl Tls12ClientSessionValue {
166    pub(crate) fn new(
167        suite: &'static Tls12CipherSuite,
168        session_id: SessionId,
169        ticket: Arc<PayloadU16>,
170        master_secret: &[u8],
171        server_cert_chain: CertificateChain<'static>,
172        server_cert_verifier: &Arc<dyn ServerCertVerifier>,
173        client_creds: &Arc<dyn ResolvesClientCert>,
174        time_now: UnixTime,
175        lifetime_secs: u32,
176        extended_ms: bool,
177    ) -> Self {
178        Self {
179            suite,
180            session_id,
181            extended_ms,
182            common: ClientSessionCommon::new(
183                ticket,
184                master_secret,
185                time_now,
186                lifetime_secs,
187                server_cert_chain,
188                server_cert_verifier,
189                client_creds,
190            ),
191        }
192    }
193
194    pub(crate) fn ticket(&mut self) -> Arc<PayloadU16> {
195        self.common.ticket.clone()
196    }
197
198    pub(crate) fn extended_ms(&self) -> bool {
199        self.extended_ms
200    }
201
202    pub(crate) fn suite(&self) -> &'static Tls12CipherSuite {
203        self.suite
204    }
205
206    #[doc(hidden)]
207    /// Test only: rewind epoch by `delta` seconds.
208    pub fn rewind_epoch(&mut self, delta: u32) {
209        self.common.epoch -= delta as u64;
210    }
211}
212
213#[cfg(feature = "tls12")]
214impl core::ops::Deref for Tls12ClientSessionValue {
215    type Target = ClientSessionCommon;
216
217    fn deref(&self) -> &Self::Target {
218        &self.common
219    }
220}
221
222#[derive(Debug, Clone)]
223pub struct ClientSessionCommon {
224    ticket: Arc<PayloadU16>,
225    secret: Zeroizing<PayloadU8>,
226    epoch: u64,
227    lifetime_secs: u32,
228    server_cert_chain: Arc<CertificateChain<'static>>,
229    server_cert_verifier: Weak<dyn ServerCertVerifier>,
230    client_creds: Weak<dyn ResolvesClientCert>,
231}
232
233impl ClientSessionCommon {
234    fn new(
235        ticket: Arc<PayloadU16>,
236        secret: &[u8],
237        time_now: UnixTime,
238        lifetime_secs: u32,
239        server_cert_chain: CertificateChain<'static>,
240        server_cert_verifier: &Arc<dyn ServerCertVerifier>,
241        client_creds: &Arc<dyn ResolvesClientCert>,
242    ) -> Self {
243        Self {
244            ticket,
245            secret: Zeroizing::new(PayloadU8::new(secret.to_vec())),
246            epoch: time_now.as_secs(),
247            lifetime_secs: cmp::min(lifetime_secs, MAX_TICKET_LIFETIME),
248            server_cert_chain: Arc::new(server_cert_chain),
249            server_cert_verifier: Arc::downgrade(server_cert_verifier),
250            client_creds: Arc::downgrade(client_creds),
251        }
252    }
253
254    pub(crate) fn compatible_config(
255        &self,
256        server_cert_verifier: &Arc<dyn ServerCertVerifier>,
257        client_creds: &Arc<dyn ResolvesClientCert>,
258    ) -> bool {
259        let same_verifier = Weak::ptr_eq(
260            &Arc::downgrade(server_cert_verifier),
261            &self.server_cert_verifier,
262        );
263        let same_creds = Weak::ptr_eq(&Arc::downgrade(client_creds), &self.client_creds);
264
265        match (same_verifier, same_creds) {
266            (true, true) => true,
267            (false, _) => {
268                crate::log::trace!("resumption not allowed between different ServerCertVerifiers");
269                false
270            }
271            (_, _) => {
272                crate::log::trace!(
273                    "resumption not allowed between different ResolvesClientCert values"
274                );
275                false
276            }
277        }
278    }
279
280    pub(crate) fn server_cert_chain(&self) -> &CertificateChain<'static> {
281        &self.server_cert_chain
282    }
283
284    pub(crate) fn secret(&self) -> &[u8] {
285        self.secret.0.as_ref()
286    }
287
288    pub(crate) fn ticket(&self) -> &[u8] {
289        self.ticket.0.as_ref()
290    }
291}
292
293static MAX_TICKET_LIFETIME: u32 = 7 * 24 * 60 * 60;
294
295/// This is the maximum allowed skew between server and client clocks, over
296/// the maximum ticket lifetime period.  This encompasses TCP retransmission
297/// times in case packet loss occurs when the client sends the ClientHello
298/// or receives the NewSessionTicket, _and_ actual clock skew over this period.
299static MAX_FRESHNESS_SKEW_MS: u32 = 60 * 1000;
300
301// --- Server types ---
302#[derive(Debug)]
303pub struct ServerSessionValue {
304    pub(crate) sni: Option<DnsName<'static>>,
305    pub(crate) version: ProtocolVersion,
306    pub(crate) cipher_suite: CipherSuite,
307    pub(crate) master_secret: Zeroizing<PayloadU8>,
308    pub(crate) extended_ms: bool,
309    pub(crate) client_cert_chain: Option<CertificateChain<'static>>,
310    pub(crate) alpn: Option<PayloadU8>,
311    pub(crate) application_data: PayloadU16,
312    pub creation_time_sec: u64,
313    pub(crate) age_obfuscation_offset: u32,
314    freshness: Option<bool>,
315}
316
317impl Codec<'_> for ServerSessionValue {
318    fn encode(&self, bytes: &mut Vec<u8>) {
319        if let Some(sni) = &self.sni {
320            1u8.encode(bytes);
321            let sni_bytes: &str = sni.as_ref();
322            PayloadU8::<MaybeEmpty>::encode_slice(sni_bytes.as_bytes(), bytes);
323        } else {
324            0u8.encode(bytes);
325        }
326        self.version.encode(bytes);
327        self.cipher_suite.encode(bytes);
328        self.master_secret.encode(bytes);
329        (u8::from(self.extended_ms)).encode(bytes);
330        if let Some(chain) = &self.client_cert_chain {
331            1u8.encode(bytes);
332            chain.encode(bytes);
333        } else {
334            0u8.encode(bytes);
335        }
336        if let Some(alpn) = &self.alpn {
337            1u8.encode(bytes);
338            alpn.encode(bytes);
339        } else {
340            0u8.encode(bytes);
341        }
342        self.application_data.encode(bytes);
343        self.creation_time_sec.encode(bytes);
344        self.age_obfuscation_offset
345            .encode(bytes);
346    }
347
348    fn read(r: &mut Reader<'_>) -> Result<Self, InvalidMessage> {
349        let has_sni = u8::read(r)?;
350        let sni = if has_sni == 1 {
351            let dns_name = PayloadU8::<MaybeEmpty>::read(r)?;
352            let dns_name = match DnsName::try_from(dns_name.0.as_slice()) {
353                Ok(dns_name) => dns_name.to_owned(),
354                Err(_) => return Err(InvalidMessage::InvalidServerName),
355            };
356
357            Some(dns_name)
358        } else {
359            None
360        };
361
362        let v = ProtocolVersion::read(r)?;
363        let cs = CipherSuite::read(r)?;
364        let ms = Zeroizing::new(PayloadU8::read(r)?);
365        let ems = u8::read(r)?;
366        let has_ccert = u8::read(r)? == 1;
367        let ccert = if has_ccert {
368            Some(CertificateChain::read(r)?.into_owned())
369        } else {
370            None
371        };
372        let has_alpn = u8::read(r)? == 1;
373        let alpn = if has_alpn {
374            Some(PayloadU8::read(r)?)
375        } else {
376            None
377        };
378        let application_data = PayloadU16::read(r)?;
379        let creation_time_sec = u64::read(r)?;
380        let age_obfuscation_offset = u32::read(r)?;
381
382        Ok(Self {
383            sni,
384            version: v,
385            cipher_suite: cs,
386            master_secret: ms,
387            extended_ms: ems == 1u8,
388            client_cert_chain: ccert,
389            alpn,
390            application_data,
391            creation_time_sec,
392            age_obfuscation_offset,
393            freshness: None,
394        })
395    }
396}
397
398impl ServerSessionValue {
399    pub(crate) fn new(
400        sni: Option<&DnsName<'_>>,
401        v: ProtocolVersion,
402        cs: CipherSuite,
403        ms: &[u8],
404        client_cert_chain: Option<CertificateChain<'static>>,
405        alpn: Option<ProtocolName>,
406        application_data: Vec<u8>,
407        creation_time: UnixTime,
408        age_obfuscation_offset: u32,
409    ) -> Self {
410        Self {
411            sni: sni.map(|dns| dns.to_owned()),
412            version: v,
413            cipher_suite: cs,
414            master_secret: Zeroizing::new(PayloadU8::new(ms.to_vec())),
415            extended_ms: false,
416            client_cert_chain,
417            alpn: alpn.map(|p| PayloadU8::new(p.as_ref().to_vec())),
418            application_data: PayloadU16::new(application_data),
419            creation_time_sec: creation_time.as_secs(),
420            age_obfuscation_offset,
421            freshness: None,
422        }
423    }
424
425    #[cfg(feature = "tls12")]
426    pub(crate) fn set_extended_ms_used(&mut self) {
427        self.extended_ms = true;
428    }
429
430    pub(crate) fn set_freshness(
431        mut self,
432        obfuscated_client_age_ms: u32,
433        time_now: UnixTime,
434    ) -> Self {
435        let client_age_ms = obfuscated_client_age_ms.wrapping_sub(self.age_obfuscation_offset);
436        let server_age_ms = (time_now
437            .as_secs()
438            .saturating_sub(self.creation_time_sec) as u32)
439            .saturating_mul(1000);
440
441        let age_difference = server_age_ms.abs_diff(client_age_ms);
442
443        self.freshness = Some(age_difference <= MAX_FRESHNESS_SKEW_MS);
444        self
445    }
446
447    pub(crate) fn is_fresh(&self) -> bool {
448        self.freshness.unwrap_or_default()
449    }
450}
451
452#[cfg(test)]
453mod tests {
454    use super::*;
455
456    #[cfg(feature = "std")] // for UnixTime::now
457    #[test]
458    fn serversessionvalue_is_debug() {
459        use std::{println, vec};
460        let ssv = ServerSessionValue::new(
461            None,
462            ProtocolVersion::TLSv1_3,
463            CipherSuite::TLS13_AES_128_GCM_SHA256,
464            &[1, 2, 3],
465            None,
466            None,
467            vec![4, 5, 6],
468            UnixTime::now(),
469            0x12345678,
470        );
471        println!("{ssv:?}");
472    }
473
474    #[test]
475    fn serversessionvalue_no_sni() {
476        let bytes = [
477            0x00, 0x03, 0x03, 0xc0, 0x23, 0x03, 0x01, 0x02, 0x03, 0x00, 0x00, 0x00, 0x00, 0x00,
478            0x12, 0x23, 0x34, 0x45, 0x56, 0x67, 0x78, 0x89, 0xfe, 0xed, 0xf0, 0x0d,
479        ];
480        let mut rd = Reader::init(&bytes);
481        let ssv = ServerSessionValue::read(&mut rd).unwrap();
482        assert_eq!(ssv.get_encoding(), bytes);
483    }
484
485    #[test]
486    fn serversessionvalue_with_cert() {
487        let bytes = [
488            0x00, 0x03, 0x03, 0xc0, 0x23, 0x03, 0x01, 0x02, 0x03, 0x00, 0x00, 0x00, 0x00, 0x00,
489            0x12, 0x23, 0x34, 0x45, 0x56, 0x67, 0x78, 0x89, 0xfe, 0xed, 0xf0, 0x0d,
490        ];
491        let mut rd = Reader::init(&bytes);
492        let ssv = ServerSessionValue::read(&mut rd).unwrap();
493        assert_eq!(ssv.get_encoding(), bytes);
494    }
495}