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
16pub const DEFAULT_SESSION_TTL: Duration = Duration::from_secs(24 * 60 * 60);
20pub const MAX_SESSION_TTL: Duration = DEFAULT_SESSION_TTL;
21
22#[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#[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 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 pub fn reset(&self) -> io::Result<()> {
104 self.primitive.reset()
105 }
106
107 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}