1use std::collections::{BTreeMap, HashMap};
37use std::str::FromStr;
38use std::sync::atomic::{AtomicBool, Ordering};
39use std::sync::Arc;
40use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
41
42use mongreldb_protocol::prepared::{PreparedStatementBinding, StatementId};
43use mongreldb_protocol::request::{AuthenticatedIdentity, IsolationLevel, SessionId};
44use mongreldb_protocol::session::{Session, TransactionState};
45use mongreldb_query::MongrelSession;
46use mongreldb_types::hlc::HlcTimestamp;
47use mongreldb_types::ids::{DatabaseId, TransactionId};
48
49pub struct SessionEntry {
51 session: std::sync::RwLock<Arc<MongrelSession>>,
56 pub owner: String,
59 last_used: std::sync::Mutex<Instant>,
60 pub lock: tokio::sync::Mutex<()>,
62 closed: AtomicBool,
67 record: std::sync::Mutex<Session>,
72 prepared_names: std::sync::Mutex<BTreeMap<String, StatementId>>,
76}
77
78impl SessionEntry {
79 pub fn session(&self) -> Arc<MongrelSession> {
82 self.session
83 .read()
84 .unwrap_or_else(|error| error.into_inner())
85 .clone()
86 }
87
88 pub(crate) fn replace_session(&self, session: MongrelSession) {
95 *self
96 .session
97 .write()
98 .unwrap_or_else(|error| error.into_inner()) = Arc::new(session);
99 }
100
101 pub(crate) fn touch(&self) {
102 if let Ok(mut t) = self.last_used.lock() {
103 *t = Instant::now();
104 }
105 if let Ok(mut record) = self.record.lock() {
106 record.last_activity_unix_micros = now_unix_micros();
107 }
108 }
109
110 pub(crate) fn is_closed(&self) -> bool {
112 self.closed.load(Ordering::Acquire)
113 }
114
115 fn mark_closed(&self) {
116 self.closed.store(true, Ordering::Release);
117 }
118
119 fn idle_for_at_least(&self, timeout: Duration) -> bool {
120 self.last_used
121 .lock()
122 .map(|t| t.elapsed() >= timeout)
123 .unwrap_or(false)
124 }
125
126 pub fn protocol_record(&self) -> Session {
129 self.record
130 .lock()
131 .map(|record| record.clone())
132 .unwrap_or_else(|error| error.into_inner().clone())
133 }
134
135 pub(crate) fn sync_record_after_request(&self, commit: Option<HlcTimestamp>) {
153 let Ok(mut record) = self.record.lock() else {
154 return;
155 };
156 record.last_activity_unix_micros = now_unix_micros();
157 let staging = self.session().staged_sql_operation_count().is_some();
158 match (&record.transaction_state, staging) {
159 (TransactionState::Idle, true) => {
160 record.transaction_state = TransactionState::Active {
161 transaction_id: TransactionId::new_random(),
162 isolation: IsolationLevel::Snapshot,
163 };
164 }
165 (TransactionState::Active { .. }, false) => {
166 record.transaction_state = TransactionState::Idle;
167 }
168 _ => {}
169 }
170 if let Some(commit_ts) = commit {
171 record.read_your_writes_token = Some(commit_ts);
172 }
173 }
174
175 pub(crate) fn allocate_statement_id(&self) -> StatementId {
178 let record = self
179 .record
180 .lock()
181 .unwrap_or_else(|error| error.into_inner());
182 let next = record
183 .prepared_statements
184 .keys()
185 .next_back()
186 .map_or(1, |id| id.get().saturating_add(1));
187 StatementId::new(next)
188 }
189
190 pub(crate) fn prepared_binding(
193 &self,
194 name: &str,
195 ) -> Option<(StatementId, PreparedStatementBinding)> {
196 let names = self
197 .prepared_names
198 .lock()
199 .unwrap_or_else(|error| error.into_inner());
200 let statement_id = names.get(name).copied()?;
201 let record = self
202 .record
203 .lock()
204 .unwrap_or_else(|error| error.into_inner());
205 record
206 .prepared_statements
207 .get(&statement_id)
208 .cloned()
209 .map(|binding| (statement_id, binding))
210 }
211
212 #[allow(dead_code)] pub(crate) fn prepared_binding_by_id(
214 &self,
215 statement_id: u64,
216 ) -> Option<(String, PreparedStatementBinding)> {
217 let statement_id = StatementId::new(statement_id);
218 let name = self
219 .prepared_names
220 .lock()
221 .unwrap_or_else(|error| error.into_inner())
222 .iter()
223 .find_map(|(name, id)| (*id == statement_id).then(|| name.clone()))?;
224 let binding = self
225 .record
226 .lock()
227 .unwrap_or_else(|error| error.into_inner())
228 .prepared_statements
229 .get(&statement_id)
230 .cloned()?;
231 Some((name, binding))
232 }
233
234 pub(crate) fn insert_prepared_binding(&self, name: String, binding: PreparedStatementBinding) {
237 self.prepared_names
238 .lock()
239 .unwrap_or_else(|error| error.into_inner())
240 .insert(name, binding.statement_id);
241 self.record
242 .lock()
243 .unwrap_or_else(|error| error.into_inner())
244 .prepared_statements
245 .insert(binding.statement_id, binding);
246 }
247
248 pub(crate) fn remove_prepared_binding(&self, name: &str) -> Option<StatementId> {
250 let statement_id = self
251 .prepared_names
252 .lock()
253 .unwrap_or_else(|error| error.into_inner())
254 .remove(name)?;
255 self.record
256 .lock()
257 .unwrap_or_else(|error| error.into_inner())
258 .prepared_statements
259 .remove(&statement_id);
260 Some(statement_id)
261 }
262
263 #[cfg(test)]
265 pub(crate) fn prepared_statement_count(&self) -> usize {
266 self.record
267 .lock()
268 .map(|record| record.prepared_statements.len())
269 .unwrap_or(0)
270 }
271}
272
273pub(crate) fn now_unix_micros() -> u64 {
276 SystemTime::now()
277 .duration_since(UNIX_EPOCH)
278 .unwrap_or_default()
279 .as_micros()
280 .min(u128::from(u64::MAX)) as u64
281}
282
283pub struct SessionStore {
287 sessions: std::sync::Mutex<HashMap<String, Arc<SessionEntry>>>,
288 max_sessions: usize,
289 idle_timeout: Duration,
290 database_id: DatabaseId,
295}
296
297impl SessionStore {
298 pub fn new(max_sessions: usize, idle_timeout: Duration) -> Self {
301 Self::new_with_database_id(max_sessions, idle_timeout, DatabaseId::new_random())
302 }
303
304 pub fn new_with_database_id(
307 max_sessions: usize,
308 idle_timeout: Duration,
309 database_id: DatabaseId,
310 ) -> Self {
311 Self {
312 sessions: std::sync::Mutex::new(HashMap::new()),
313 max_sessions: max_sessions.max(1),
316 idle_timeout,
317 database_id,
318 }
319 }
320
321 pub fn database_id(&self) -> DatabaseId {
323 self.database_id
324 }
325
326 pub fn create(&self, session: MongrelSession, owner: String) -> Option<String> {
332 self.create_with_identity(session, owner, AuthenticatedIdentity::Credentialless)
333 }
334
335 pub fn create_with_identity(
338 &self,
339 session: MongrelSession,
340 owner: String,
341 principal: AuthenticatedIdentity,
342 ) -> Option<String> {
343 let mut guard = self.sessions.lock().ok()?;
344 if guard.len() >= self.max_sessions {
345 return None;
346 }
347 let token = random_token()?;
348 let session_id = SessionId::from_str(&token).unwrap_or(SessionId::ZERO);
351 let record = Session::new(session_id, principal, self.database_id, now_unix_micros());
352 guard.insert(
353 token.clone(),
354 Arc::new(SessionEntry {
355 session: std::sync::RwLock::new(Arc::new(session)),
356 owner,
357 last_used: std::sync::Mutex::new(Instant::now()),
358 lock: tokio::sync::Mutex::new(()),
359 closed: AtomicBool::new(false),
360 record: std::sync::Mutex::new(record),
361 prepared_names: std::sync::Mutex::new(BTreeMap::new()),
362 }),
363 );
364 Some(token)
365 }
366
367 pub fn get(&self, token: &str, owner: &str) -> Option<Arc<SessionEntry>> {
374 let guard = self.sessions.lock().ok()?;
375 let entry = guard.get(token)?;
376 if entry.owner != owner || entry.is_closed() {
377 return None;
378 }
379 Some(Arc::clone(entry))
380 }
381
382 #[allow(dead_code)] pub(crate) fn get_by_token(&self, token: &str) -> Option<Arc<SessionEntry>> {
387 let guard = self.sessions.lock().ok()?;
388 let entry = guard.get(token)?;
389 (!entry.is_closed()).then(|| Arc::clone(entry))
390 }
391
392 #[allow(dead_code)] pub(crate) fn close_by_token(&self, token: &str) -> bool {
394 let Ok(mut guard) = self.sessions.lock() else {
395 return false;
396 };
397 if let Some(entry) = guard.remove(token) {
398 entry.mark_closed();
399 true
400 } else {
401 false
402 }
403 }
404
405 pub fn close(&self, token: &str, owner: &str) -> bool {
409 self.take_for_close(token, owner).is_some()
410 }
411
412 pub(crate) fn take_for_close(&self, token: &str, owner: &str) -> Option<Arc<SessionEntry>> {
415 if let Ok(mut guard) = self.sessions.lock() {
416 if let Some(entry) = guard.get(token) {
417 if entry.owner != owner {
418 return None;
419 }
420 entry.mark_closed();
421 }
422 return guard.remove(token);
423 }
424 None
425 }
426
427 pub fn sweep_idle(&self) -> usize {
431 let Ok(mut guard) = self.sessions.lock() else {
432 return 0;
433 };
434 let timeout = self.idle_timeout;
435 let to_evict: Vec<String> = guard
440 .iter()
441 .filter_map(|(token, entry)| {
442 if entry.idle_for_at_least(timeout)
443 && entry.session().query_registry().active_for_session(token) == 0
444 && entry.lock.try_lock().is_ok()
445 {
446 Some(token.clone())
447 } else {
448 None
449 }
450 })
451 .collect();
452 let count = to_evict.len();
453 for token in &to_evict {
454 if let Some(entry) = guard.remove(token) {
455 entry.mark_closed();
456 }
457 }
458 count
459 }
460
461 pub fn len(&self) -> usize {
463 self.sessions.lock().map(|g| g.len()).unwrap_or(0)
464 }
465
466 pub fn is_empty(&self) -> bool {
467 self.len() == 0
468 }
469
470 pub(crate) fn close_all(&self) {
472 if let Ok(mut guard) = self.sessions.lock() {
473 for entry in guard.values() {
474 entry.mark_closed();
475 }
476 guard.clear();
477 }
478 }
479}
480
481pub fn spawn_session_reaper(store: Arc<SessionStore>) {
485 std::thread::Builder::new()
486 .name("mongreldb-session-reaper".into())
487 .spawn(move || loop {
488 std::thread::sleep(Duration::from_secs(30));
489 let evicted = store.sweep_idle();
490 if evicted > 0 {
491 eprintln!("[session-reaper] evicted {evicted} idle session(s)");
492 }
493 })
494 .expect("spawn session-reaper");
495}
496
497fn random_token() -> Option<String> {
503 let bytes = read_urandom(16)?;
504 let mut n = 0u128;
505 for &b in &bytes {
506 n = (n << 8) | b as u128;
507 }
508 Some(format!("{n:032x}"))
509}
510
511fn read_urandom(n: usize) -> Option<Vec<u8>> {
512 use std::io::Read;
513 let mut f = std::fs::File::open("/dev/urandom").ok()?;
514 let mut buf = vec![0u8; n];
515 f.read_exact(&mut buf).ok()?;
516 Some(buf)
517}
518
519#[cfg(test)]
520mod tests {
521 use super::*;
522 use mongreldb_core::Database;
523 use mongreldb_query::{RegisteredQueryGuard, SqlQueryOptions};
524 use tempfile::tempdir;
525
526 fn make_session() -> MongrelSession {
527 let dir = tempdir().unwrap();
528 let db = Arc::new(Database::create(dir.path()).unwrap());
529 std::mem::forget(dir);
532 MongrelSession::open(db).unwrap()
533 }
534
535 #[test]
536 fn create_and_get_roundtrip() {
537 let store = SessionStore::new(8, Duration::from_secs(60));
538 let token = store.create(make_session(), "alice".into()).unwrap();
539 assert!(store.get(&token, "alice").is_some());
540 assert!(store.get(&token, "eve").is_none());
542 assert_eq!(store.len(), 1);
543 }
544
545 #[test]
546 fn close_removes_session() {
547 let store = SessionStore::new(8, Duration::from_secs(60));
548 let token = store.create(make_session(), "alice".into()).unwrap();
549 assert!(store.close(&token, "alice"));
550 assert!(store.get(&token, "alice").is_none());
551 assert!(store.is_empty());
552 let t2 = store.create(make_session(), "bob".into()).unwrap();
554 assert!(!store.close(&t2, "alice"));
555 assert_eq!(store.len(), 1);
556 }
557
558 #[test]
559 fn capacity_limit_rejects_new_sessions() {
560 let store = SessionStore::new(1, Duration::from_secs(60));
561 assert!(store.create(make_session(), "a".into()).is_some());
562 assert!(store.create(make_session(), "b".into()).is_none());
564 assert_eq!(store.len(), 1);
565 }
566
567 #[test]
568 fn sweep_idle_evicts_stale_sessions() {
569 let store = SessionStore::new(8, Duration::from_millis(1));
570 let token = store.create(make_session(), "alice".into()).unwrap();
571 assert_eq!(store.len(), 1);
572 std::thread::sleep(Duration::from_millis(20));
574 let evicted = store.sweep_idle();
575 assert_eq!(evicted, 1);
576 assert!(store.get(&token, "alice").is_none());
577 assert!(store.is_empty());
578 }
579
580 #[test]
581 fn sweep_idle_keeps_active_queries() {
582 let store = SessionStore::new(8, Duration::from_millis(1));
583 let token = store.create(make_session(), "alice".into()).unwrap();
584 let entry = store.get(&token, "alice").unwrap();
585 let query = entry
586 .session()
587 .register_query(SqlQueryOptions {
588 session_id: Some(token.clone()),
589 ..SqlQueryOptions::default()
590 })
591 .unwrap();
592 let query = RegisteredQueryGuard::new(query);
593 std::thread::sleep(Duration::from_millis(20));
594
595 assert_eq!(store.sweep_idle(), 0);
596 assert!(store.get(&token, "alice").is_some());
597
598 drop(query);
599 assert_eq!(store.sweep_idle(), 1);
600 assert!(store.is_empty());
601 }
602
603 #[test]
604 fn protocol_record_carries_identity_database_and_session_id() {
605 let database_id = DatabaseId::new_random();
606 let store = SessionStore::new_with_database_id(8, Duration::from_secs(60), database_id);
607 let identity = AuthenticatedIdentity::CatalogUser {
608 username: "alice".to_owned(),
609 user_id: 42,
610 created_version: 7,
611 };
612 let token = store
613 .create_with_identity(make_session(), "alice".into(), identity.clone())
614 .unwrap();
615 let entry = store.get(&token, "alice").unwrap();
616 let record = entry.protocol_record();
617 assert_eq!(record.session_id, SessionId::from_str(&token).unwrap());
618 assert_eq!(record.principal, identity);
619 assert_eq!(record.current_database, database_id);
620 assert_eq!(record.transaction_state, TransactionState::Idle);
621 assert!(record.prepared_statements.is_empty());
622 assert!(record.settings.is_empty());
623 assert_eq!(record.read_your_writes_token, None);
624 assert!(record.last_activity_unix_micros > 0);
625 assert_eq!(store.database_id(), database_id);
626 }
627
628 #[test]
629 fn sync_record_tracks_commit_and_activity() {
630 let store = SessionStore::new(8, Duration::from_secs(60));
631 let token = store.create(make_session(), "alice".into()).unwrap();
632 let entry = store.get(&token, "alice").unwrap();
633 let before = entry.protocol_record().last_activity_unix_micros;
634 std::thread::sleep(Duration::from_millis(2));
635
636 let commit_ts = HlcTimestamp {
637 physical_micros: now_unix_micros().saturating_sub(1_000),
638 logical: 3,
639 node_tiebreaker: 0,
640 };
641 entry.sync_record_after_request(Some(commit_ts));
642 let record = entry.protocol_record();
643 assert!(record.last_activity_unix_micros >= before);
644 assert_eq!(
645 record.read_your_writes_token,
646 Some(commit_ts),
647 "a committed request must advance the read-your-writes token to the commit timestamp"
648 );
649
650 assert_eq!(record.transaction_state, TransactionState::Idle);
652
653 entry.sync_record_after_request(None);
655 assert_eq!(
656 entry.protocol_record().read_your_writes_token,
657 Some(commit_ts)
658 );
659 }
660
661 #[test]
662 fn prepared_bindings_are_tracked_by_name_and_id() {
663 let store = SessionStore::new(8, Duration::from_secs(60));
664 let token = store.create(make_session(), "alice".into()).unwrap();
665 let entry = store.get(&token, "alice").unwrap();
666
667 let first = entry.allocate_statement_id();
668 assert_eq!(first, StatementId::new(1));
669 let mut binding = PreparedStatementBinding {
670 statement_id: first,
671 sql: "SELECT 1".to_owned(),
672 parameter_types: vec![],
673 catalog_version: mongreldb_types::ids::MetadataVersion::new(3),
674 schema_versions: BTreeMap::new(),
675 feature_set: Default::default(),
676 };
677 entry.insert_prepared_binding("stmt_a".to_owned(), binding.clone());
678 assert_eq!(entry.prepared_statement_count(), 1);
679 assert_eq!(
680 entry.prepared_binding("stmt_a"),
681 Some((first, binding.clone()))
682 );
683
684 let second = entry.allocate_statement_id();
686 assert_eq!(second, StatementId::new(2));
687 binding.statement_id = second;
688 entry.insert_prepared_binding("stmt_b".to_owned(), binding);
689
690 assert_eq!(entry.remove_prepared_binding("stmt_a"), Some(first));
691 assert_eq!(entry.prepared_binding("stmt_a"), None);
692 assert_eq!(entry.prepared_statement_count(), 1);
693 assert_eq!(entry.protocol_record().prepared_statements.len(), 1);
694 assert_eq!(entry.remove_prepared_binding("stmt_a"), None);
696 }
697}