use std::collections::HashMap as StdHashMap;
use std::sync::{Arc, Mutex as StdMutex};
use std::time::Duration;
use surrealdb_core::dbs::Session;
use surrealdb_core::iam::{Auth, Level};
use surrealdb_core::kvs::Datastore;
use surrealdb_core::rpc::{
DbResult, Method, RpcProtocol, method_not_allowed, method_not_found, session_exists,
session_not_found, types_error_from_anyhow,
};
use surrealdb_types::{Array, Error as TypesError, HashMap, Value};
use tokio::sync::{Mutex as AsyncMutex, OwnedMutexGuard, RwLock};
use uuid::Uuid;
use web_time::{SystemTime, UNIX_EPOCH};
use crate::cnf::{HTTP_MAX_ATTACHED_SESSIONS, PKG_NAME, PKG_VERSION};
fn now_ms() -> u64 {
match SystemTime::now().duration_since(UNIX_EPOCH) {
Ok(duration) => duration.as_millis() as u64,
Err(error) => panic!("Clock may have gone backwards: {:?}", error.duration()),
}
}
pub struct Http {
kvs: Arc<Datastore>,
sessions: HashMap<Uuid, Arc<RwLock<Session>>>,
ephemeral_sessions: HashMap<Uuid, ()>,
durable_session_ttl: Option<Duration>,
session_locks: SessionLocks,
durable_deadlines: HashMap<Uuid, u64>,
}
impl Http {
pub fn new(kvs: Arc<Datastore>) -> Self {
Self::new_with_durability(kvs, None)
}
pub fn new_with_durability(kvs: Arc<Datastore>, durable_session_ttl: Option<Duration>) -> Self {
Self {
kvs,
sessions: HashMap::new(),
ephemeral_sessions: HashMap::new(),
durable_session_ttl,
session_locks: SessionLocks::default(),
durable_deadlines: HashMap::new(),
}
}
pub(crate) fn register_ephemeral_session(&self, id: Uuid, session: Arc<RwLock<Session>>) {
self.ephemeral_sessions.insert(id, ());
self.sessions.insert(id, session);
}
pub(crate) fn remove_ephemeral_session(&self, id: &Uuid) {
if self.ephemeral_sessions.contains_key(id) {
self.sessions.remove(id);
self.ephemeral_sessions.remove(id);
}
}
fn attached_session_count(&self) -> usize {
self.sessions.len().saturating_sub(self.ephemeral_sessions.len())
}
fn prune_expired_cached_sessions(&self) {
let now = now_ms();
for (id, deadline) in self.durable_deadlines.to_vec() {
if deadline <= now {
self.session_map().remove(&id);
self.durable_deadlines.remove(&id);
}
}
}
pub(crate) async fn verify_caller_for_session(
&self,
session_id: &Uuid,
caller_au: &Auth,
) -> Result<(), TypesError> {
if self.ephemeral_sessions.contains_key(session_id) {
return Err(session_not_found(*session_id));
}
let session_lock = self.get_session(session_id).await?;
let session_guard = session_lock.read().await;
let session_au = session_guard.au.as_ref();
if caller_may_use_session(session_au, caller_au) {
Ok(())
} else {
Err(session_not_found(*session_id))
}
}
pub(crate) fn session_locks(&self) -> &SessionLocks {
&self.session_locks
}
pub(crate) async fn revalidate_cached_session(&self, id: &Uuid) {
if self.durable_session_ttl.is_none() {
return;
}
if self.ephemeral_sessions.contains_key(id) {
return;
}
let Some(lock) = self.sessions.get(id) else {
return;
};
match self.kvs.load_rpc_session(*id).await {
Ok(Some((stored, expires_at))) => {
*lock.write().await = stored;
self.durable_deadlines.insert(*id, expires_at);
}
Ok(None) => {
self.session_map().remove(id);
self.durable_deadlines.remove(id);
}
Err(err) => {
warn!("Failed to revalidate durable RPC session {id}: {err}");
}
}
}
pub(crate) async fn touch_durable_session(&self, id: &Uuid) {
let Some(ttl) = self.durable_session_ttl else {
return;
};
if self.ephemeral_sessions.contains_key(id) {
return;
}
let refresh = match self.durable_deadlines.get(id) {
Some(deadline) => deadline.saturating_sub(now_ms()) < ttl.as_millis() as u64 / 2,
None => true,
};
if !refresh {
return;
}
let Some(lock) = self.sessions.get(id) else {
return;
};
let session = lock.read().await.clone();
match self.kvs.update_rpc_session(*id, &session, ttl).await {
Ok(true) => {
self.durable_deadlines.insert(*id, now_ms() + ttl.as_millis() as u64);
}
Ok(false) => {
self.session_map().remove(id);
self.durable_deadlines.remove(id);
}
Err(err) => {
error!("Failed to refresh durable RPC session {id}: {err}");
}
}
}
}
type SessionLockMap = StdMutex<StdHashMap<Uuid, Arc<AsyncMutex<()>>>>;
#[derive(Default)]
pub(crate) struct SessionLocks {
locks: SessionLockMap,
}
impl SessionLocks {
pub(crate) async fn acquire(&self, id: Uuid) -> SessionDispatchGuard<'_> {
let mutex = {
let mut locks = self.locks.lock().expect("session lock map poisoned");
Arc::clone(locks.entry(id).or_default())
};
let mut pending = PendingAcquire {
locks: &self.locks,
id,
armed: true,
};
let permit = mutex.lock_owned().await;
pending.armed = false;
SessionDispatchGuard {
locks: &self.locks,
id,
permit: Some(permit),
}
}
#[cfg(test)]
fn len(&self) -> usize {
self.locks.lock().expect("session lock map poisoned").len()
}
}
fn prune_idle_lock(locks: &SessionLockMap, id: &Uuid) {
let Ok(mut map) = locks.lock() else {
return;
};
if let Some(mutex) = map.get(id)
&& Arc::strong_count(mutex) == 1
{
map.remove(id);
}
}
struct PendingAcquire<'a> {
locks: &'a SessionLockMap,
id: Uuid,
armed: bool,
}
impl Drop for PendingAcquire<'_> {
fn drop(&mut self) {
if self.armed {
prune_idle_lock(self.locks, &self.id);
}
}
}
pub(crate) struct SessionDispatchGuard<'a> {
locks: &'a SessionLockMap,
id: Uuid,
permit: Option<OwnedMutexGuard<()>>,
}
impl Drop for SessionDispatchGuard<'_> {
fn drop(&mut self) {
self.permit.take();
prune_idle_lock(self.locks, &self.id);
}
}
fn caller_may_use_session(session_au: &Auth, caller_au: &Auth) -> bool {
match session_au.level() {
Level::No => true,
_ => session_au.id() == caller_au.id() && session_au.level() == caller_au.level(),
}
}
impl RpcProtocol for Http {
fn kvs(&self) -> &Datastore {
&self.kvs
}
fn kvs_arc(&self) -> Arc<Datastore> {
Arc::clone(&self.kvs)
}
fn version_data(&self) -> DbResult {
let value = Value::String(format!("{PKG_NAME}-{}", *PKG_VERSION));
DbResult::Other(value)
}
fn session_map(&self) -> &HashMap<Uuid, Arc<RwLock<Session>>> {
&self.sessions
}
async fn sessions(&self) -> Result<DbResult, TypesError> {
Err(method_not_allowed(Method::Sessions.to_string()))
}
async fn attach(&self, session_id: Uuid) -> Result<DbResult, TypesError> {
if self.session_map().contains_key(&session_id) {
return Err(session_exists(session_id));
}
if self.attached_session_count() >= *HTTP_MAX_ATTACHED_SESSIONS {
self.prune_expired_cached_sessions();
if self.attached_session_count() >= *HTTP_MAX_ATTACHED_SESSIONS {
return Err(method_not_allowed(Method::Attach.to_string()));
}
}
let mut session = Session::default().with_rt(Self::LQ_SUPPORT);
session.id = Some(session_id);
if let Some(ttl) = self.durable_session_ttl {
let created = self
.kvs
.create_rpc_session(session_id, &session, ttl)
.await
.map_err(types_error_from_anyhow)?;
if !created {
return Err(session_exists(session_id));
}
self.durable_deadlines.insert(session_id, now_ms() + ttl.as_millis() as u64);
}
self.session_map().insert(session_id, Arc::new(RwLock::new(session)));
Ok(DbResult::Other(Value::None))
}
const PERSIST_SESSIONS: bool = true;
fn persist_sessions_enabled(&self) -> bool {
self.durable_session_ttl.is_some()
}
async fn load_session(&self, id: &Uuid) -> Option<Session> {
self.durable_session_ttl?;
if self.ephemeral_sessions.contains_key(id) {
return None;
}
if self.attached_session_count() >= *HTTP_MAX_ATTACHED_SESSIONS {
self.prune_expired_cached_sessions();
if self.attached_session_count() >= *HTTP_MAX_ATTACHED_SESSIONS {
warn!(
"Refusing to rehydrate RPC session {id}: the attached session limit is reached"
);
return None;
}
}
match self.kvs.load_rpc_session(*id).await {
Ok(Some((session, expires_at))) => {
self.durable_deadlines.insert(*id, expires_at);
Some(session)
}
Ok(None) => None,
Err(err) => {
warn!("Failed to load durable RPC session {id}: {err}");
None
}
}
}
async fn persist_session(&self, id: &Uuid, session: &Session) {
let Some(ttl) = self.durable_session_ttl else {
return;
};
if self.ephemeral_sessions.contains_key(id) {
return;
}
match self.kvs.update_rpc_session(*id, session, ttl).await {
Ok(true) => {
self.durable_deadlines.insert(*id, now_ms() + ttl.as_millis() as u64);
}
Ok(false) => {
self.session_map().remove(id);
self.durable_deadlines.remove(id);
}
Err(err) => {
error!("Failed to persist durable RPC session {id}: {err}");
}
}
}
async fn forget_session(&self, id: &Uuid) -> Result<(), TypesError> {
if self.durable_session_ttl.is_none() {
return Ok(());
}
self.kvs.delete_rpc_session(*id).await.map_err(types_error_from_anyhow)?;
self.durable_deadlines.remove(id);
Ok(())
}
const LQ_SUPPORT: bool = false;
async fn cleanup_lqs(&self, _session_id: &Uuid) {
}
async fn cleanup_all_lqs(&self) {
}
async fn begin(&self, _txn: Option<Uuid>, _session_id: Uuid) -> Result<DbResult, TypesError> {
Err(method_not_found(Method::Begin.to_string()))
}
async fn commit(
&self,
_txn: Option<Uuid>,
_session_id: Uuid,
_params: Array,
) -> Result<DbResult, TypesError> {
Err(method_not_found(Method::Commit.to_string()))
}
async fn cancel(
&self,
_txn: Option<Uuid>,
_session_id: Uuid,
_params: Array,
) -> Result<DbResult, TypesError> {
Err(method_not_found(Method::Cancel.to_string()))
}
}
#[cfg(test)]
mod tests {
use surrealdb_core::dbs::Capabilities;
use surrealdb_core::iam::Role;
use super::*;
const TTL: Duration = Duration::from_secs(3600);
async fn mem_ds() -> Arc<Datastore> {
Arc::new(
Datastore::builder()
.with_capabilities(Capabilities::all())
.build_with_path("memory")
.await
.unwrap(),
)
}
fn http(ds: &Arc<Datastore>) -> Http {
Http::new_with_durability(Arc::clone(ds), Some(TTL))
}
fn use_params(ns: &str, db: &str) -> Array {
vec![Value::String(ns.to_owned()), Value::String(db.to_owned())].into()
}
async fn attach_owner_session(rpc: &Http, sid: Uuid, ns: &str, db: &str) {
rpc.execute(None, sid, Some(sid), Method::Attach, Array::new()).await.unwrap();
{
let lock = rpc.get_session(&sid).await.unwrap();
lock.write().await.au = Arc::new(Auth::for_root(Role::Owner));
}
rpc.execute(None, sid, Some(sid), Method::Use, use_params(ns, db)).await.unwrap();
}
#[tokio::test]
async fn attached_session_survives_a_new_instance() {
let ds = mem_ds().await;
let sid = Uuid::new_v4();
let node_a = http(&ds);
attach_owner_session(&node_a, sid, "app", "app").await;
let node_b = http(&ds);
{
let lock = node_b.get_session(&sid).await.expect("session was not rehydrated");
let session = lock.read().await;
assert!(session.au.is_root(), "auth lost on rehydrate");
assert_eq!(session.ns.as_deref(), Some("app"));
assert_eq!(session.db.as_deref(), Some("app"));
}
node_b
.execute(None, sid, Some(sid), Method::Use, use_params("app", "other"))
.await
.unwrap();
let lock = node_b.get_session(&sid).await.unwrap();
assert_eq!(lock.read().await.db.as_deref(), Some("other"));
}
#[tokio::test]
async fn caller_gate_applies_to_rehydrated_sessions() {
let ds = mem_ds().await;
let sid = Uuid::new_v4();
attach_owner_session(&http(&ds), sid, "app", "app").await;
let node_b = http(&ds);
assert!(node_b.verify_caller_for_session(&sid, &Auth::default()).await.is_err());
assert!(node_b.verify_caller_for_session(&sid, &Auth::for_root(Role::Owner)).await.is_ok());
}
#[tokio::test]
async fn detach_prevents_rehydration_everywhere() {
let ds = mem_ds().await;
let sid = Uuid::new_v4();
let node_a = http(&ds);
attach_owner_session(&node_a, sid, "app", "app").await;
assert!(ds.load_rpc_session(sid).await.unwrap().is_some());
node_a.execute(None, sid, Some(sid), Method::Detach, Array::new()).await.unwrap();
assert!(ds.load_rpc_session(sid).await.unwrap().is_none());
let node_b = http(&ds);
assert!(node_b.get_session(&sid).await.is_err());
}
#[tokio::test]
async fn disabled_instances_never_touch_the_datastore() {
let ds = mem_ds().await;
let sid = Uuid::new_v4();
let off = Http::new(Arc::clone(&ds));
assert!(!off.persist_sessions_enabled());
off.execute(None, sid, Some(sid), Method::Attach, Array::new()).await.unwrap();
off.execute(None, sid, Some(sid), Method::Use, use_params("app", "app")).await.unwrap();
assert!(ds.load_rpc_session(sid).await.unwrap().is_none());
assert!(http(&ds).get_session(&sid).await.is_err());
}
#[tokio::test]
async fn ephemeral_sessions_are_never_persisted() {
let ds = mem_ds().await;
let rpc = http(&ds);
let rid = Uuid::new_v4();
rpc.register_ephemeral_session(rid, Arc::new(RwLock::new(Session::owner())));
rpc.execute(None, rid, None, Method::Use, use_params("app", "app")).await.unwrap();
assert!(ds.load_rpc_session(rid).await.unwrap().is_none());
rpc.persist_session(&rid, &Session::owner()).await;
assert!(ds.load_rpc_session(rid).await.unwrap().is_none());
ds.persist_rpc_session(rid, &Session::owner(), TTL).await.unwrap();
assert!(rpc.load_session(&rid).await.is_none());
}
#[tokio::test]
async fn rehydration_respects_the_attached_session_cap() {
let ds = mem_ds().await;
let rpc = http(&ds);
for _ in 0..*HTTP_MAX_ATTACHED_SESSIONS {
rpc.set_session(Uuid::new_v4(), Arc::new(RwLock::new(Session::default())));
}
let sid = Uuid::new_v4();
ds.persist_rpc_session(sid, &Session::owner(), TTL).await.unwrap();
assert!(rpc.load_session(&sid).await.is_none());
assert!(ds.load_rpc_session(sid).await.unwrap().is_some());
}
#[tokio::test]
async fn prune_reclaims_durably_expired_cached_sessions() {
let ds = mem_ds().await;
let rpc = http(&ds);
let live = Uuid::new_v4();
let expired = Uuid::new_v4();
rpc.set_session(live, Arc::new(RwLock::new(Session::owner())));
rpc.durable_deadlines.insert(live, now_ms() + 60_000);
rpc.set_session(expired, Arc::new(RwLock::new(Session::owner())));
rpc.durable_deadlines.insert(expired, now_ms().saturating_sub(1));
rpc.prune_expired_cached_sessions();
assert!(rpc.session_map().contains_key(&live), "a live cached session is kept");
assert!(
!rpc.session_map().contains_key(&expired),
"a durably-expired cached session is pruned to reclaim cap space"
);
assert!(rpc.durable_deadlines.get(&expired).is_none(), "its deadline is dropped too");
}
#[tokio::test]
async fn rehydration_does_not_refresh_the_durable_ttl() {
let ds = mem_ds().await;
let sid = Uuid::new_v4();
attach_owner_session(&http(&ds), sid, "app", "app").await;
let (_, before) =
ds.load_rpc_session(sid).await.unwrap().expect("precondition: session persisted");
let node_b = http(&ds);
assert!(node_b.load_session(&sid).await.is_some(), "session should rehydrate");
let (_, after) = ds.load_rpc_session(sid).await.unwrap().expect("session still present");
assert_eq!(after, before, "rehydration must not rewrite the durable expiry");
assert_eq!(
node_b.durable_deadlines.get(&sid),
Some(before),
"rehydration should cache the stored expiry unchanged"
);
}
#[tokio::test]
async fn cached_session_is_evicted_when_the_durable_copy_expires() {
let ds = mem_ds().await;
let rpc = http(&ds);
let sid = Uuid::new_v4();
attach_owner_session(&rpc, sid, "app", "app").await;
assert!(rpc.session_map().contains_key(&sid));
rpc.revalidate_cached_session(&sid).await;
assert!(rpc.session_map().contains_key(&sid), "a live cached session must not be evicted");
ds.persist_rpc_session(sid, &Session::owner(), Duration::from_millis(1)).await.unwrap();
tokio::time::sleep(Duration::from_millis(5)).await;
rpc.revalidate_cached_session(&sid).await;
assert!(
!rpc.session_map().contains_key(&sid),
"a cached session whose durable copy expired must be evicted"
);
assert!(rpc.durable_deadlines.get(&sid).is_none(), "its deadline must be dropped too");
}
#[tokio::test]
async fn touch_refreshes_only_when_the_ttl_runs_low() {
let ds = mem_ds().await;
let rpc = http(&ds);
let sid = Uuid::new_v4();
rpc.execute(None, sid, Some(sid), Method::Attach, Array::new()).await.unwrap();
let fresh = rpc.durable_deadlines.get(&sid).expect("attach should record a deadline");
rpc.touch_durable_session(&sid).await;
assert_eq!(rpc.durable_deadlines.get(&sid), Some(fresh));
let aged = now_ms() + TTL.as_millis() as u64 / 4;
rpc.durable_deadlines.insert(sid, aged);
rpc.touch_durable_session(&sid).await;
let refreshed =
rpc.durable_deadlines.get(&sid).expect("touch should keep the deadline entry");
assert!(refreshed > aged, "touch did not refresh the durable deadline");
}
#[tokio::test]
async fn read_only_request_observes_a_remote_detach() {
let ds = mem_ds().await;
let rpc = http(&ds);
let sid = Uuid::new_v4();
attach_owner_session(&rpc, sid, "app", "app").await;
assert!(rpc.session_map().contains_key(&sid));
ds.delete_rpc_session(sid).await.unwrap();
rpc.revalidate_cached_session(&sid).await;
assert!(
!rpc.session_map().contains_key(&sid),
"a remotely-detached session must be evicted before dispatch"
);
}
#[tokio::test]
async fn revalidation_adopts_a_remote_mutation_that_keeps_the_entry() {
let ds = mem_ds().await;
let sid = Uuid::new_v4();
attach_owner_session(&http(&ds), sid, "app", "app").await;
let node_b = http(&ds);
{
let lock = node_b.get_session(&sid).await.expect("B rehydrates");
assert!(lock.read().await.au.is_root(), "B caches the authenticated session");
}
let cleared = Session {
id: Some(sid),
..Default::default()
};
ds.persist_rpc_session(sid, &cleared, TTL).await.unwrap();
node_b.revalidate_cached_session(&sid).await;
let lock = node_b.get_session(&sid).await.expect("still cached");
assert!(!lock.read().await.au.is_root(), "B must adopt the auth cleared on another node");
}
#[tokio::test]
async fn create_rpc_session_is_conditional() {
let ds = mem_ds().await;
let sid = Uuid::new_v4();
assert!(
ds.create_rpc_session(sid, &Session::owner(), TTL).await.unwrap(),
"the first create must succeed"
);
assert!(
!ds.create_rpc_session(sid, &Session::owner(), TTL).await.unwrap(),
"a create for an existing id must be refused, not overwrite"
);
}
#[tokio::test]
async fn attach_reports_session_exists_for_a_durable_copy_created_elsewhere() {
let ds = mem_ds().await;
let sid = Uuid::new_v4();
ds.create_rpc_session(sid, &Session::owner(), TTL).await.unwrap();
let rpc = http(&ds);
assert!(
rpc.attach(sid).await.is_err(),
"attach must not overwrite a durable session created on another node"
);
}
#[tokio::test]
async fn attach_creates_durable_but_refresh_never_resurrects() {
let ds = mem_ds().await;
let rpc = http(&ds);
let sid = Uuid::new_v4();
rpc.execute(None, sid, Some(sid), Method::Attach, Array::new()).await.unwrap();
assert!(
ds.load_rpc_session(sid).await.unwrap().is_some(),
"attach must create the durable copy"
);
ds.delete_rpc_session(sid).await.unwrap();
rpc.persist_session(&sid, &Session::owner()).await;
assert!(
ds.load_rpc_session(sid).await.unwrap().is_none(),
"a refresh must not resurrect a session deleted on another node"
);
assert!(
!rpc.session_map().contains_key(&sid),
"the revoked session must be dropped from the local cache"
);
}
#[tokio::test]
async fn session_lock_slot_is_freed_after_a_canceled_waiter() {
let locks = SessionLocks::default();
let id = Uuid::new_v4();
let held = locks.acquire(id).await;
assert_eq!(locks.len(), 1);
let waiter = tokio::time::timeout(Duration::from_millis(20), locks.acquire(id)).await;
assert!(waiter.is_err(), "the waiter should time out while the slot is held");
assert_eq!(locks.len(), 1, "the holder still owns the slot");
drop(held);
assert_eq!(locks.len(), 0, "a released slot with no waiters must be pruned");
}
#[tokio::test]
async fn session_lock_slot_is_freed_when_a_registered_waiter_is_dropped() {
let locks = SessionLocks::default();
let id = Uuid::new_v4();
let held = locks.acquire(id).await;
let mut waiter = Box::pin(locks.acquire(id));
tokio::select! {
biased;
_ = &mut waiter => panic!("the waiter must not acquire while the slot is held"),
_ = std::future::ready(()) => {}
}
assert_eq!(locks.len(), 1);
drop(held);
drop(waiter);
assert_eq!(locks.len(), 0, "a canceled registered waiter must not orphan its slot");
}
}