Skip to main content

orbit_rustls/session/
mod.rs

1use std::sync::Arc;
2use std::sync::atomic::{AtomicU64, Ordering};
3use std::time::Duration;
4use std::{fmt, io};
5
6use orbit_core::Fleet;
7
8use self::store::{DOMAIN_MAX, SessionPrimitive};
9
10mod store;
11#[cfg(feature = "rustls_0_23")]
12mod v23;
13#[cfg(feature = "rustls_0_24")]
14mod v24;
15
16/// Default retention for opaque rustls server-session values.
17///
18/// This is no longer than rustls 0.23's stateful TLS 1.3 ticket lifetime.
19pub const DEFAULT_SESSION_TTL: Duration = Duration::from_secs(24 * 60 * 60);
20pub const MAX_SESSION_TTL: Duration = DEFAULT_SESSION_TTL;
21
22/// Stable isolation boundary for one equivalent set of rustls server configs.
23///
24/// Transport family, authentication policy, and any compatibility epoch must
25/// be reflected in this value. Public TLS and mTLS must not share a domain.
26#[derive(Clone, PartialEq, Eq, Hash)]
27pub struct SessionDomain(Arc<[u8]>);
28
29impl SessionDomain {
30    pub fn new(value: impl AsRef<[u8]>) -> io::Result<Self> {
31        let value = value.as_ref();
32        if value.is_empty() {
33            return Err(io::Error::new(
34                io::ErrorKind::InvalidInput,
35                "rustls session domain must not be empty",
36            ));
37        }
38        if value.len() > DOMAIN_MAX {
39            return Err(io::Error::new(
40                io::ErrorKind::InvalidInput,
41                format!(
42                    "rustls session domain is too long: {} bytes (maximum {DOMAIN_MAX})",
43                    value.len()
44                ),
45            ));
46        }
47        Ok(Self(Arc::from(value)))
48    }
49
50    fn as_bytes(&self) -> &[u8] {
51        &self.0
52    }
53}
54
55impl fmt::Debug for SessionDomain {
56    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
57        formatter
58            .debug_tuple("SessionDomain")
59            .field(&String::from_utf8_lossy(&self.0))
60            .finish()
61    }
62}
63
64/// One fleet-wide session table from which isolated rustls views are made.
65#[derive(Clone)]
66pub struct FleetServerSessions {
67    primitive: Arc<SessionPrimitive>,
68    ttl_ms: Arc<AtomicU64>,
69}
70
71impl FleetServerSessions {
72    pub fn open(fleet: Arc<Fleet>) -> io::Result<Self> {
73        Self::with_ttl(fleet, DEFAULT_SESSION_TTL)
74    }
75
76    pub fn with_ttl(fleet: Arc<Fleet>, ttl: Duration) -> io::Result<Self> {
77        let ttl_ms = validate_ttl(ttl)?;
78        Ok(Self {
79            primitive: Arc::new(SessionPrimitive::open(&fleet)?),
80            ttl_ms: Arc::new(AtomicU64::new(ttl_ms)),
81        })
82    }
83
84    pub fn set_ttl(&self, ttl: Duration) -> io::Result<()> {
85        self.ttl_ms.store(validate_ttl(ttl)?, Ordering::Release);
86        Ok(())
87    }
88
89    pub fn ttl(&self) -> Duration {
90        Duration::from_millis(self.ttl_ms.load(Ordering::Acquire))
91    }
92
93    /// Create a rustls storage view isolated by `domain`.
94    pub fn storage(&self, domain: SessionDomain) -> Arc<OrbitSessionStorage> {
95        Arc::new(OrbitSessionStorage {
96            primitive: self.primitive.clone(),
97            domain,
98            ttl_ms: self.ttl_ms.clone(),
99        })
100    }
101
102    /// Clear every domain. Call only during owner-controlled quiescent boot.
103    pub fn reset(&self) -> io::Result<()> {
104        self.primitive.reset()
105    }
106
107    /// Remove the SHM name and companion lock file.
108    ///
109    /// Existing mappings remain valid. This is intended for fleet teardown
110    /// and tests, not runtime cache invalidation.
111    pub fn unlink(&self) -> io::Result<()> {
112        self.primitive.unlink()
113    }
114}
115
116impl fmt::Debug for FleetServerSessions {
117    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
118        formatter
119            .debug_struct("FleetServerSessions")
120            .field("ttl", &self.ttl())
121            .finish_non_exhaustive()
122    }
123}
124
125pub struct OrbitSessionStorage {
126    primitive: Arc<SessionPrimitive>,
127    domain: SessionDomain,
128    ttl_ms: Arc<AtomicU64>,
129}
130
131impl fmt::Debug for OrbitSessionStorage {
132    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
133        formatter
134            .debug_struct("OrbitSessionStorage")
135            .field("domain", &self.domain)
136            .field(
137                "ttl",
138                &Duration::from_millis(self.ttl_ms.load(Ordering::Acquire)),
139            )
140            .finish_non_exhaustive()
141    }
142}
143
144impl OrbitSessionStorage {
145    fn put_bytes(&self, key: &[u8], value: &[u8]) -> bool {
146        let ttl = Duration::from_millis(self.ttl_ms.load(Ordering::Acquire));
147        match self.primitive.put(self.domain.as_bytes(), key, value, ttl) {
148            Ok(cached) => {
149                tracing::trace!(
150                    target: "orbit_rustls::session_cache",
151                    operation = "put",
152                    domain = ?self.domain,
153                    key_len = key.len(),
154                    value_len = value.len(),
155                    ?ttl,
156                    cached,
157                    "rustls session cache write"
158                );
159                cached
160            }
161            Err(error) => {
162                tracing::trace!(
163                    target: "orbit_rustls::session_cache",
164                    operation = "put",
165                    domain = ?self.domain,
166                    key_len = key.len(),
167                    value_len = value.len(),
168                    ?ttl,
169                    %error,
170                    "rustls session cache write failed"
171                );
172                false
173            }
174        }
175    }
176
177    fn get_bytes(&self, key: &[u8]) -> Option<Vec<u8>> {
178        match self.primitive.get(self.domain.as_bytes(), key) {
179            Ok(value) => {
180                tracing::trace!(
181                    target: "orbit_rustls::session_cache",
182                    operation = "get",
183                    domain = ?self.domain,
184                    key_len = key.len(),
185                    value_len = value.as_ref().map_or(0, Vec::len),
186                    hit = value.is_some(),
187                    "rustls session cache read"
188                );
189                value
190            }
191            Err(error) => {
192                tracing::trace!(
193                    target: "orbit_rustls::session_cache",
194                    operation = "get",
195                    domain = ?self.domain,
196                    key_len = key.len(),
197                    %error,
198                    "rustls session cache read failed"
199                );
200                None
201            }
202        }
203    }
204
205    fn take_bytes(&self, key: &[u8]) -> Option<Vec<u8>> {
206        match self.primitive.take(self.domain.as_bytes(), key) {
207            Ok(value) => {
208                tracing::trace!(
209                    target: "orbit_rustls::session_cache",
210                    operation = "take",
211                    domain = ?self.domain,
212                    key_len = key.len(),
213                    value_len = value.as_ref().map_or(0, Vec::len),
214                    hit = value.is_some(),
215                    "rustls session cache read and consume"
216                );
217                value
218            }
219            Err(error) => {
220                tracing::trace!(
221                    target: "orbit_rustls::session_cache",
222                    operation = "take",
223                    domain = ?self.domain,
224                    key_len = key.len(),
225                    %error,
226                    "rustls session cache read and consume failed"
227                );
228                None
229            }
230        }
231    }
232}
233
234fn validate_ttl(ttl: Duration) -> io::Result<u64> {
235    if ttl.is_zero() {
236        return Err(io::Error::new(
237            io::ErrorKind::InvalidInput,
238            "rustls session TTL must be greater than zero",
239        ));
240    }
241    if ttl > MAX_SESSION_TTL {
242        return Err(io::Error::new(
243            io::ErrorKind::InvalidInput,
244            format!(
245                "rustls session TTL {ttl:?} exceeds rustls' stateful lifetime {MAX_SESSION_TTL:?}"
246            ),
247        ));
248    }
249    u64::try_from(ttl.as_millis().max(1)).map_err(|_| {
250        io::Error::new(
251            io::ErrorKind::InvalidInput,
252            "rustls session TTL does not fit milliseconds",
253        )
254    })
255}
256
257#[cfg(test)]
258mod tests {
259    use super::*;
260
261    fn sessions() -> FleetServerSessions {
262        let fleet = Arc::new(Fleet::join("rustls-session-unit", 1).expect("fleet"));
263        FleetServerSessions::open(fleet).expect("sessions")
264    }
265
266    #[test]
267    fn rustls_view_put_get_and_take() {
268        let sessions = sessions();
269        let storage = sessions.storage(SessionDomain::new("tcp-public").expect("domain"));
270
271        assert!(storage.put_bytes(b"ticket", b"secret"));
272        assert_eq!(storage.get_bytes(b"ticket"), Some(b"secret".to_vec()));
273        assert_eq!(storage.take_bytes(b"ticket"), Some(b"secret".to_vec()));
274        assert_eq!(storage.take_bytes(b"ticket"), None);
275    }
276
277    #[test]
278    fn domains_are_isolated() {
279        let sessions = sessions();
280        let tcp = sessions.storage(SessionDomain::new("tcp-public").expect("domain"));
281        let quic = sessions.storage(SessionDomain::new("quic-public").expect("domain"));
282
283        assert!(tcp.put_bytes(b"same-ticket", b"tcp"));
284        assert_eq!(quic.take_bytes(b"same-ticket"), None);
285        assert_eq!(tcp.take_bytes(b"same-ticket"), Some(b"tcp".to_vec()));
286    }
287
288    #[test]
289    fn existing_views_observe_ttl_updates() {
290        let sessions = sessions();
291        let storage = sessions.storage(SessionDomain::new("tcp-public").expect("domain"));
292        sessions
293            .set_ttl(Duration::from_millis(1))
294            .expect("TTL update");
295
296        assert!(storage.put_bytes(b"short", b"secret"));
297        std::thread::sleep(Duration::from_millis(5));
298        assert_eq!(storage.get_bytes(b"short"), None);
299    }
300
301    #[test]
302    fn ttl_cannot_outlive_rustls_stateful_tickets() {
303        let fleet = Arc::new(Fleet::join("rustls-session-ttl", 1).expect("fleet"));
304
305        assert!(
306            FleetServerSessions::with_ttl(fleet, MAX_SESSION_TTL + Duration::from_millis(1))
307                .is_err()
308        );
309    }
310}