contextvm_sdk/transport/server/
session_store.rs1use std::num::NonZeroUsize;
11use std::sync::Arc;
12
13use lru::LruCache;
14use tokio::sync::RwLock;
15
16use crate::core::types::ClientSession;
17use crate::transport::server::ServerEventRouteStore;
18
19const LOG_TARGET: &str = "contextvm_sdk::transport::server::session_store";
20
21pub const DEFAULT_MAX_SESSIONS: usize = 1000;
26
27pub type EvictionCallback = Arc<dyn Fn(String) + Send + Sync>;
30
31#[derive(Clone)]
35pub struct SessionStore {
36 sessions: Arc<RwLock<LruCache<String, ClientSession>>>,
37 on_evicted: Option<EvictionCallback>,
38}
39
40impl Default for SessionStore {
41 fn default() -> Self {
42 Self::new()
43 }
44}
45
46impl SessionStore {
47 pub fn new() -> Self {
49 Self::with_capacity(DEFAULT_MAX_SESSIONS)
50 }
51
52 pub fn with_capacity(max_sessions: usize) -> Self {
54 Self {
55 sessions: Arc::new(RwLock::new(LruCache::new(
56 NonZeroUsize::new(max_sessions).unwrap_or(NonZeroUsize::new(1).unwrap()),
57 ))),
58 on_evicted: None,
59 }
60 }
61
62 pub fn set_eviction_callback(&mut self, cb: EvictionCallback) {
64 self.on_evicted = Some(cb);
65 }
66
67 pub fn eviction_callback(&self) -> Option<EvictionCallback> {
69 self.on_evicted.clone()
70 }
71
72 pub async fn get_or_create_session(
78 &self,
79 client_pubkey: &str,
80 is_encrypted: bool,
81 event_routes: &ServerEventRouteStore,
82 ) -> bool {
83 let on_evicted = self.on_evicted.clone();
84 let mut sessions = self.sessions.write().await;
85 if let Some(session) = sessions.get_mut(client_pubkey) {
86 session.is_encrypted = is_encrypted;
87 false
88 } else {
89 let new_session = ClientSession::new(is_encrypted);
90 let evicted = sessions.push(client_pubkey.to_string(), new_session);
91 Self::handle_eviction(
92 client_pubkey,
93 evicted,
94 &mut sessions,
95 on_evicted.as_ref(),
96 event_routes,
97 )
98 .await;
99 true
100 }
101 }
102
103 pub async fn get_session(&self, client_pubkey: &str) -> Option<SessionSnapshot> {
106 let sessions = self.sessions.read().await;
107 sessions.peek(client_pubkey).map(|s| SessionSnapshot {
108 is_initialized: s.is_initialized,
109 is_encrypted: s.is_encrypted,
110 has_sent_common_tags: s.has_sent_common_tags,
111 supports_ephemeral_gift_wrap: s.supports_ephemeral_gift_wrap,
112 supports_encryption: s.supports_encryption,
113 supports_ephemeral_encryption: s.supports_ephemeral_encryption,
114 supports_oversized_transfer: s.supports_oversized_transfer,
115 supports_open_stream: s.supports_open_stream,
116 })
117 }
118
119 pub async fn mark_initialized(&self, client_pubkey: &str) -> bool {
121 let mut sessions = self.sessions.write().await;
122 if let Some(session) = sessions.get_mut(client_pubkey) {
123 session.is_initialized = true;
124 true
125 } else {
126 false
127 }
128 }
129
130 pub async fn mark_common_tags_sent(&self, client_pubkey: &str) -> bool {
132 let mut sessions = self.sessions.write().await;
133 if let Some(session) = sessions.get_mut(client_pubkey) {
134 session.has_sent_common_tags = true;
135 true
136 } else {
137 false
138 }
139 }
140
141 pub async fn remove_session(&self, client_pubkey: &str) -> bool {
143 self.sessions.write().await.pop(client_pubkey).is_some()
144 }
145
146 pub async fn clear(&self) {
148 self.sessions.write().await.clear();
149 }
150
151 pub async fn session_count(&self) -> usize {
153 self.sessions.read().await.len()
154 }
155
156 pub async fn get_all_sessions(&self) -> Vec<(String, SessionSnapshot)> {
158 let sessions = self.sessions.read().await;
159 sessions
160 .iter()
161 .map(|(k, s)| {
162 (
163 k.clone(),
164 SessionSnapshot {
165 is_initialized: s.is_initialized,
166 is_encrypted: s.is_encrypted,
167 has_sent_common_tags: s.has_sent_common_tags,
168 supports_ephemeral_gift_wrap: s.supports_ephemeral_gift_wrap,
169 supports_encryption: s.supports_encryption,
170 supports_ephemeral_encryption: s.supports_ephemeral_encryption,
171 supports_oversized_transfer: s.supports_oversized_transfer,
172 supports_open_stream: s.supports_open_stream,
173 },
174 )
175 })
176 .collect()
177 }
178
179 pub(crate) async fn write(
181 &self,
182 ) -> tokio::sync::RwLockWriteGuard<'_, LruCache<String, ClientSession>> {
183 self.sessions.write().await
184 }
185
186 pub(crate) async fn read(
188 &self,
189 ) -> tokio::sync::RwLockReadGuard<'_, LruCache<String, ClientSession>> {
190 self.sessions.read().await
191 }
192
193 pub(crate) async fn handle_eviction(
200 inserted_key: &str,
201 evicted: Option<(String, ClientSession)>,
202 sessions: &mut LruCache<String, ClientSession>,
203 on_evicted: Option<&EvictionCallback>,
204 event_routes: &ServerEventRouteStore,
205 ) {
206 if let Some((evicted_key, evicted_session)) = evicted {
207 if evicted_key != inserted_key {
210 if event_routes
211 .has_active_routes_for_client(&evicted_key)
212 .await
213 {
214 tracing::warn!(
215 target: LOG_TARGET,
216 client_pubkey = %evicted_key,
217 "LRU eviction of session with active routes; recreating with clean state"
218 );
219 let _ = sessions.push(
223 evicted_key.clone(),
224 ClientSession::new(evicted_session.is_encrypted),
225 );
226 } else if let Some(cb) = on_evicted {
227 cb(evicted_key);
228 }
229 }
230 }
231 }
232}
233
234#[derive(Debug, Clone, PartialEq, Eq)]
237pub struct SessionSnapshot {
238 pub is_initialized: bool,
240 pub is_encrypted: bool,
242 pub has_sent_common_tags: bool,
244 pub supports_ephemeral_gift_wrap: bool,
246 pub supports_encryption: bool,
248 pub supports_ephemeral_encryption: bool,
250 pub supports_oversized_transfer: bool,
252 pub supports_open_stream: bool,
254}
255
256#[cfg(test)]
257mod tests {
258 use super::*;
259 use serde_json::json;
260
261 fn routes() -> ServerEventRouteStore {
262 ServerEventRouteStore::new()
263 }
264
265 #[tokio::test]
266 async fn create_and_retrieve_session() {
267 let store = SessionStore::new();
268 let r = routes();
269
270 let created = store.get_or_create_session("client-1", true, &r).await;
271 assert!(created);
272
273 let snap = store.get_session("client-1").await.unwrap();
274 assert!(snap.is_encrypted);
275 assert!(!snap.is_initialized);
276 }
277
278 #[tokio::test]
279 async fn get_or_create_returns_existing() {
280 let store = SessionStore::new();
281 let r = routes();
282
283 let created = store.get_or_create_session("client-1", false, &r).await;
284 assert!(created);
285
286 let created2 = store.get_or_create_session("client-1", true, &r).await;
287 assert!(!created2);
288
289 let snap = store.get_session("client-1").await.unwrap();
290 assert!(snap.is_encrypted);
291 }
292
293 #[tokio::test]
294 async fn mark_initialized() {
295 let store = SessionStore::new();
296 let r = routes();
297 store.get_or_create_session("client-1", false, &r).await;
298
299 assert!(store.mark_initialized("client-1").await);
300 let snap = store.get_session("client-1").await.unwrap();
301 assert!(snap.is_initialized);
302 }
303
304 #[tokio::test]
305 async fn mark_initialized_unknown_returns_false() {
306 let store = SessionStore::new();
307 assert!(!store.mark_initialized("unknown").await);
308 }
309
310 #[tokio::test]
311 async fn remove_session() {
312 let store = SessionStore::new();
313 let r = routes();
314 store.get_or_create_session("client-1", false, &r).await;
315 assert!(store.remove_session("client-1").await);
316 assert!(store.get_session("client-1").await.is_none());
317 }
318
319 #[tokio::test]
320 async fn remove_unknown_returns_false() {
321 let store = SessionStore::new();
322 assert!(!store.remove_session("unknown").await);
323 }
324
325 #[tokio::test]
326 async fn clear_all_sessions() {
327 let store = SessionStore::new();
328 let r = routes();
329 store.get_or_create_session("client-1", false, &r).await;
330 store.get_or_create_session("client-2", true, &r).await;
331
332 store.clear().await;
333
334 assert_eq!(store.session_count().await, 0);
335 assert!(store.get_session("client-1").await.is_none());
336 assert!(store.get_session("client-2").await.is_none());
337 }
338
339 #[tokio::test]
340 async fn get_all_sessions() {
341 let store = SessionStore::new();
342 let r = routes();
343 store.get_or_create_session("client-1", false, &r).await;
344 store.get_or_create_session("client-2", true, &r).await;
345
346 let all = store.get_all_sessions().await;
347 assert_eq!(all.len(), 2);
348
349 let keys: Vec<&str> = all.iter().map(|(k, _)| k.as_str()).collect();
350 assert!(keys.contains(&"client-1"));
351 assert!(keys.contains(&"client-2"));
352 }
353
354 #[tokio::test]
357 async fn new_session_capability_fields_default_false() {
358 let store = SessionStore::new();
359 let r = routes();
360 store.get_or_create_session("client-1", false, &r).await;
361
362 let sessions = store.read().await;
363 let session = sessions.peek("client-1").unwrap();
364 assert!(!session.has_sent_common_tags);
365 assert!(!session.supports_encryption);
366 assert!(!session.supports_ephemeral_encryption);
367 assert!(!session.supports_oversized_transfer);
368 }
369
370 #[tokio::test]
371 async fn snapshot_surfaces_learned_capabilities() {
372 let store = SessionStore::new();
373 let r = routes();
374 store.get_or_create_session("client-1", false, &r).await;
375
376 let snap = store.get_session("client-1").await.unwrap();
378 assert!(!snap.supports_encryption);
379 assert!(!snap.supports_ephemeral_encryption);
380 assert!(!snap.supports_oversized_transfer);
381 assert!(!snap.supports_open_stream);
382
383 {
385 let mut sessions = store.write().await;
386 let session = sessions.get_mut("client-1").unwrap();
387 session.supports_encryption = true;
388 session.supports_ephemeral_encryption = true;
389 session.supports_oversized_transfer = true;
390 session.supports_open_stream = true;
391 }
392
393 let snap = store.get_session("client-1").await.unwrap();
394 assert!(snap.supports_encryption);
395 assert!(snap.supports_ephemeral_encryption);
396 assert!(snap.supports_oversized_transfer);
397 assert!(snap.supports_open_stream);
398
399 let all = store.get_all_sessions().await;
401 let (_, snap_all) = all.iter().find(|(k, _)| k == "client-1").unwrap();
402 assert!(snap_all.supports_encryption);
403 assert!(snap_all.supports_ephemeral_encryption);
404 assert!(snap_all.supports_oversized_transfer);
405 assert!(snap_all.supports_open_stream);
406 }
407
408 #[tokio::test]
409 async fn has_sent_common_tags_flag() {
410 let store = SessionStore::new();
411 let r = routes();
412 store.get_or_create_session("client-1", false, &r).await;
413
414 let mut sessions = store.write().await;
415 let session = sessions.get_mut("client-1").unwrap();
416 assert!(!session.has_sent_common_tags);
417 session.has_sent_common_tags = true;
418 assert!(session.has_sent_common_tags);
419 }
420
421 #[tokio::test]
422 async fn capability_or_assign_persists() {
423 let store = SessionStore::new();
424 let r = routes();
425 store.get_or_create_session("client-1", false, &r).await;
426
427 {
428 let mut sessions = store.write().await;
429 let session = sessions.get_mut("client-1").unwrap();
430 session.supports_encryption |= true;
431 session.supports_ephemeral_encryption |= false;
432 }
433
434 {
435 let mut sessions = store.write().await;
436 let session = sessions.get_mut("client-1").unwrap();
437 session.supports_encryption |= false;
438 session.supports_ephemeral_encryption |= true;
439 }
440
441 let sessions = store.read().await;
442 let session = sessions.peek("client-1").unwrap();
443 assert!(session.supports_encryption, "OR-assign must not downgrade");
444 assert!(session.supports_ephemeral_encryption);
445 assert!(!session.supports_oversized_transfer);
446 }
447
448 #[tokio::test]
449 async fn capability_fields_independent_per_client() {
450 let store = SessionStore::new();
451 let r = routes();
452 store.get_or_create_session("client-a", false, &r).await;
453 store.get_or_create_session("client-b", false, &r).await;
454
455 {
456 let mut sessions = store.write().await;
457 let sa = sessions.get_mut("client-a").unwrap();
458 sa.supports_encryption = true;
459 sa.has_sent_common_tags = true;
460 }
461
462 let sessions = store.read().await;
463 let sa = sessions.peek("client-a").unwrap();
464 let sb = sessions.peek("client-b").unwrap();
465 assert!(sa.supports_encryption);
466 assert!(sa.has_sent_common_tags);
467 assert!(!sb.supports_encryption);
468 assert!(!sb.has_sent_common_tags);
469 }
470
471 #[tokio::test]
472 async fn get_or_create_preserves_capability_fields() {
473 let store = SessionStore::new();
474 let r = routes();
475 store.get_or_create_session("client-1", false, &r).await;
476
477 {
478 let mut sessions = store.write().await;
479 let session = sessions.get_mut("client-1").unwrap();
480 session.supports_encryption = true;
481 session.has_sent_common_tags = true;
482 }
483
484 let created = store.get_or_create_session("client-1", true, &r).await;
485 assert!(!created);
486
487 let sessions = store.read().await;
488 let session = sessions.peek("client-1").unwrap();
489 assert!(session.supports_encryption);
490 assert!(session.has_sent_common_tags);
491 }
492
493 #[tokio::test]
494 async fn clear_resets_capability_fields() {
495 let store = SessionStore::new();
496 let r = routes();
497 store.get_or_create_session("client-1", false, &r).await;
498 {
499 let mut sessions = store.write().await;
500 let s = sessions.get_mut("client-1").unwrap();
501 s.supports_encryption = true;
502 }
503
504 store.clear().await;
505 store.get_or_create_session("client-1", false, &r).await;
506
507 let sessions = store.read().await;
508 let session = sessions.peek("client-1").unwrap();
509 assert!(!session.supports_encryption);
510 assert!(!session.has_sent_common_tags);
511 }
512
513 #[tokio::test]
516 async fn lru_eviction_drops_oldest_session() {
517 let store = SessionStore::with_capacity(3);
518 let r = routes();
519 store.get_or_create_session("a", false, &r).await;
520 store.get_or_create_session("b", false, &r).await;
521 store.get_or_create_session("c", false, &r).await;
522
523 store.get_or_create_session("d", false, &r).await;
524
525 assert!(
526 store.get_session("a").await.is_none(),
527 "a should be evicted"
528 );
529 assert!(store.get_session("b").await.is_some());
530 assert!(store.get_session("c").await.is_some());
531 assert!(store.get_session("d").await.is_some());
532 assert_eq!(store.session_count().await, 3);
533 }
534
535 #[tokio::test]
536 async fn eviction_callback_fires_on_lru_eviction() {
537 let evicted = Arc::new(std::sync::Mutex::new(Vec::<String>::new()));
538 let evicted_clone = evicted.clone();
539 let r = routes();
540
541 let mut store = SessionStore::with_capacity(2);
542 store.set_eviction_callback(Arc::new(move |pubkey| {
543 evicted_clone.lock().unwrap().push(pubkey);
544 }));
545
546 store.get_or_create_session("a", false, &r).await;
547 store.get_or_create_session("b", false, &r).await;
548 store.get_or_create_session("c", false, &r).await;
549
550 let evicted = evicted.lock().unwrap();
551 assert_eq!(evicted.len(), 1);
552 assert_eq!(evicted[0], "a");
553 }
554
555 #[tokio::test]
556 async fn eviction_safety_recreates_session_with_active_routes() {
557 let store = SessionStore::with_capacity(2);
558 let r = routes();
559 store.get_or_create_session("a", true, &r).await;
560 store.get_or_create_session("b", false, &r).await;
561
562 r.register("evt1".into(), "a".into(), json!(1), None).await;
564
565 store.get_or_create_session("c", false, &r).await;
568
569 let snap = store.get_session("a").await;
570 assert!(
571 snap.is_some(),
572 "session with active routes must survive eviction"
573 );
574 assert!(
576 store.get_session("b").await.is_none(),
577 "b should be evicted"
578 );
579 }
580
581 #[tokio::test]
582 async fn with_capacity_sets_limit() {
583 let store = SessionStore::with_capacity(5);
584 let r = routes();
585 for i in 0..10 {
586 store
587 .get_or_create_session(&format!("client-{i}"), false, &r)
588 .await;
589 }
590 assert_eq!(store.session_count().await, 5);
591 }
592}