use std::collections::{HashMap, VecDeque};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex, MutexGuard};
use tokio::sync::Notify;
use uuid::Uuid;
use crate::SessionError;
pub type Established<T, E = SessionError> = Result<T, E>;
const REMEMBERED_ENDED_SESSIONS: usize = 1024;
pub struct SessionEntry<T, E = SessionError> {
ready: Notify,
outcome: Mutex<Option<Established<T, E>>>,
}
impl<T, E> Default for SessionEntry<T, E> {
fn default() -> Self {
Self {
ready: Notify::default(),
outcome: Mutex::new(None),
}
}
}
impl<T: Clone, E: Clone> SessionEntry<T, E> {
pub async fn wait_ready(&self) -> Established<T, E> {
loop {
let notified = self.ready.notified();
if let Some(outcome) = self.outcome() {
return outcome;
}
notified.await;
}
}
pub fn outcome(&self) -> Option<Established<T, E>> {
self.guard().clone()
}
pub fn publish(&self, outcome: Established<T, E>) {
*self.guard() = Some(outcome);
self.ready.notify_waiters();
}
fn guard(&self) -> MutexGuard<'_, Option<Established<T, E>>> {
self.outcome.lock().expect("session registry poisoned")
}
}
pub struct SessionRegistry<T, E = SessionError> {
entries: Mutex<Entries<T, E>>,
closed: AtomicBool,
}
impl<T, E> Default for SessionRegistry<T, E> {
fn default() -> Self {
Self {
entries: Mutex::new(Entries::default()),
closed: AtomicBool::new(false),
}
}
}
struct Entries<T, E> {
live: HashMap<Uuid, Arc<SessionEntry<T, E>>>,
ended: VecDeque<Uuid>,
}
impl<T, E> Default for Entries<T, E> {
fn default() -> Self {
Self {
live: HashMap::new(),
ended: VecDeque::new(),
}
}
}
impl<T, E> Entries<T, E> {
fn end(&mut self, id: Uuid) -> Option<Arc<SessionEntry<T, E>>> {
self.ended.push_back(id);
if self.ended.len() > REMEMBERED_ENDED_SESSIONS {
self.ended.pop_front();
}
self.live.remove(&id)
}
fn has_ended(&self, id: Uuid) -> bool {
self.ended.contains(&id)
}
}
impl<T: Clone, E: Clone + From<SessionError>> SessionRegistry<T, E> {
fn map(&self) -> MutexGuard<'_, Entries<T, E>> {
self.entries.lock().expect("session registry poisoned")
}
pub async fn resolve(&self, id: Uuid) -> Established<T, E> {
let entry = {
let mut entries = self.map();
if !entries.live.contains_key(&id)
&& (entries.has_ended(id) || self.closed.load(Ordering::Acquire))
{
return Err(SessionError::NotFound(id).into());
}
Arc::clone(entries.live.entry(id).or_default())
};
entry.wait_ready().await
}
pub fn entry(&self, id: Uuid) -> Arc<SessionEntry<T, E>> {
Arc::clone(self.map().live.entry(id).or_default())
}
pub fn established(&self, id: Uuid) -> Option<Established<T, E>> {
self.map().live.get(&id).cloned()?.outcome()
}
pub fn end(&self, id: Uuid) -> Option<Established<T, E>> {
let entry = self.map().end(id)?;
let outcome = entry.outcome();
entry.publish(Err(SessionError::NotFound(id).into()));
outcome
}
pub fn close(&self) {
self.closed.store(true, Ordering::Release);
for (id, entry) in self.map().live.iter() {
if entry.outcome().is_none() {
entry.publish(Err(SessionError::NotFound(*id).into()));
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn registry() -> Arc<SessionRegistry<u8, SessionError>> {
Arc::new(SessionRegistry::default())
}
#[tokio::test]
async fn a_request_waits_for_a_session_that_is_still_being_registered() {
let sessions = registry();
let id = Uuid::new_v4();
let waiter = {
let sessions = Arc::clone(&sessions);
tokio::spawn(async move { sessions.resolve(id).await })
};
tokio::task::yield_now().await;
sessions.entry(id).publish(Ok(7));
assert_eq!(
waiter.await.unwrap().expect("registered"),
7,
"a queued registration resolves the request"
);
}
#[tokio::test]
async fn a_request_for_an_ended_session_fails() {
let sessions = registry();
let id = Uuid::new_v4();
sessions.entry(id).publish(Ok(7));
assert_eq!(sessions.resolve(id).await.expect("registered"), 7);
assert_eq!(
sessions.end(id).expect("a live session").expect("established"),
7,
"ending hands back what was established"
);
assert!(matches!(sessions.resolve(id).await, Err(SessionError::NotFound(_))));
}
#[tokio::test]
async fn ending_a_session_fails_the_request_already_waiting_on_it() {
let sessions = registry();
let id = Uuid::new_v4();
let waiter = {
let sessions = Arc::clone(&sessions);
tokio::spawn(async move { sessions.resolve(id).await })
};
tokio::task::yield_now().await;
sessions.end(id);
assert!(matches!(waiter.await.unwrap(), Err(SessionError::NotFound(_))));
}
#[tokio::test]
async fn the_reason_a_session_failed_reaches_the_request() {
let sessions: Arc<SessionRegistry<u8, surrealdb_types::Error>> =
Arc::new(SessionRegistry::default());
let id = Uuid::new_v4();
sessions.entry(id).publish(Err(surrealdb_types::Error::connection(
"the server went away".to_string(),
surrealdb_types::ConnectionError::ConnectionFailed,
)));
let error = sessions.resolve(id).await.expect_err("the session was never established");
assert!(error.is_connection(), "the failure must still read as a connection failure");
}
#[tokio::test]
async fn closing_fails_the_waiting_and_everything_after() {
let sessions = registry();
let waiting = Uuid::new_v4();
let waiter = {
let sessions = Arc::clone(&sessions);
tokio::spawn(async move { sessions.resolve(waiting).await })
};
tokio::task::yield_now().await;
sessions.close();
assert!(matches!(waiter.await.unwrap(), Err(SessionError::NotFound(_))));
assert!(matches!(sessions.resolve(Uuid::new_v4()).await, Err(SessionError::NotFound(_))));
}
}