surrealdb_engine_api/
session.rs1use std::collections::{HashMap, VecDeque};
21use std::sync::atomic::{AtomicBool, Ordering};
22use std::sync::{Arc, Mutex, MutexGuard};
23
24use tokio::sync::Notify;
25use uuid::Uuid;
26
27use crate::SessionError;
28
29pub type Established<T, E = SessionError> = Result<T, E>;
38
39const REMEMBERED_ENDED_SESSIONS: usize = 1024;
48
49pub struct SessionEntry<T, E = SessionError> {
51 ready: Notify,
52 outcome: Mutex<Option<Established<T, E>>>,
53}
54
55impl<T, E> Default for SessionEntry<T, E> {
56 fn default() -> Self {
57 Self {
58 ready: Notify::default(),
59 outcome: Mutex::new(None),
60 }
61 }
62}
63
64impl<T: Clone, E: Clone> SessionEntry<T, E> {
65 pub async fn wait_ready(&self) -> Established<T, E> {
72 loop {
73 let notified = self.ready.notified();
74 if let Some(outcome) = self.outcome() {
75 return outcome;
76 }
77 notified.await;
78 }
79 }
80
81 pub fn outcome(&self) -> Option<Established<T, E>> {
85 self.guard().clone()
86 }
87
88 pub fn publish(&self, outcome: Established<T, E>) {
93 *self.guard() = Some(outcome);
94 self.ready.notify_waiters();
95 }
96
97 fn guard(&self) -> MutexGuard<'_, Option<Established<T, E>>> {
98 self.outcome.lock().expect("session registry poisoned")
99 }
100}
101
102pub struct SessionRegistry<T, E = SessionError> {
104 entries: Mutex<Entries<T, E>>,
105 closed: AtomicBool,
111}
112
113impl<T, E> Default for SessionRegistry<T, E> {
114 fn default() -> Self {
115 Self {
116 entries: Mutex::new(Entries::default()),
117 closed: AtomicBool::new(false),
118 }
119 }
120}
121
122struct Entries<T, E> {
123 live: HashMap<Uuid, Arc<SessionEntry<T, E>>>,
124 ended: VecDeque<Uuid>,
125}
126
127impl<T, E> Default for Entries<T, E> {
128 fn default() -> Self {
129 Self {
130 live: HashMap::new(),
131 ended: VecDeque::new(),
132 }
133 }
134}
135
136impl<T, E> Entries<T, E> {
137 fn end(&mut self, id: Uuid) -> Option<Arc<SessionEntry<T, E>>> {
138 self.ended.push_back(id);
139 if self.ended.len() > REMEMBERED_ENDED_SESSIONS {
140 self.ended.pop_front();
141 }
142 self.live.remove(&id)
143 }
144
145 fn has_ended(&self, id: Uuid) -> bool {
146 self.ended.contains(&id)
147 }
148}
149
150impl<T: Clone, E: Clone + From<SessionError>> SessionRegistry<T, E> {
151 fn map(&self) -> MutexGuard<'_, Entries<T, E>> {
152 self.entries.lock().expect("session registry poisoned")
153 }
154
155 pub async fn resolve(&self, id: Uuid) -> Established<T, E> {
158 let entry = {
159 let mut entries = self.map();
160 if !entries.live.contains_key(&id)
164 && (entries.has_ended(id) || self.closed.load(Ordering::Acquire))
165 {
166 return Err(SessionError::NotFound(id).into());
167 }
168 Arc::clone(entries.live.entry(id).or_default())
169 };
170 entry.wait_ready().await
171 }
172
173 pub fn entry(&self, id: Uuid) -> Arc<SessionEntry<T, E>> {
176 Arc::clone(self.map().live.entry(id).or_default())
177 }
178
179 pub fn established(&self, id: Uuid) -> Option<Established<T, E>> {
185 self.map().live.get(&id).cloned()?.outcome()
186 }
187
188 pub fn end(&self, id: Uuid) -> Option<Established<T, E>> {
194 let entry = self.map().end(id)?;
195 let outcome = entry.outcome();
196 entry.publish(Err(SessionError::NotFound(id).into()));
197 outcome
198 }
199
200 pub fn close(&self) {
203 self.closed.store(true, Ordering::Release);
204 for (id, entry) in self.map().live.iter() {
205 if entry.outcome().is_none() {
206 entry.publish(Err(SessionError::NotFound(*id).into()));
207 }
208 }
209 }
210}
211
212#[cfg(test)]
213mod tests {
214 use super::*;
215
216 fn registry() -> Arc<SessionRegistry<u8, SessionError>> {
217 Arc::new(SessionRegistry::default())
218 }
219
220 #[tokio::test]
221 async fn a_request_waits_for_a_session_that_is_still_being_registered() {
222 let sessions = registry();
223 let id = Uuid::new_v4();
224
225 let waiter = {
226 let sessions = Arc::clone(&sessions);
227 tokio::spawn(async move { sessions.resolve(id).await })
228 };
229 tokio::task::yield_now().await;
231 sessions.entry(id).publish(Ok(7));
232
233 assert_eq!(
234 waiter.await.unwrap().expect("registered"),
235 7,
236 "a queued registration resolves the request"
237 );
238 }
239
240 #[tokio::test]
244 async fn a_request_for_an_ended_session_fails() {
245 let sessions = registry();
246 let id = Uuid::new_v4();
247
248 sessions.entry(id).publish(Ok(7));
249 assert_eq!(sessions.resolve(id).await.expect("registered"), 7);
250
251 assert_eq!(
252 sessions.end(id).expect("a live session").expect("established"),
253 7,
254 "ending hands back what was established"
255 );
256 assert!(matches!(sessions.resolve(id).await, Err(SessionError::NotFound(_))));
257 }
258
259 #[tokio::test]
262 async fn ending_a_session_fails_the_request_already_waiting_on_it() {
263 let sessions = registry();
264 let id = Uuid::new_v4();
265
266 let waiter = {
267 let sessions = Arc::clone(&sessions);
268 tokio::spawn(async move { sessions.resolve(id).await })
269 };
270 tokio::task::yield_now().await;
271 sessions.end(id);
272
273 assert!(matches!(waiter.await.unwrap(), Err(SessionError::NotFound(_))));
274 }
275
276 #[tokio::test]
281 async fn the_reason_a_session_failed_reaches_the_request() {
282 let sessions: Arc<SessionRegistry<u8, surrealdb_types::Error>> =
283 Arc::new(SessionRegistry::default());
284 let id = Uuid::new_v4();
285
286 sessions.entry(id).publish(Err(surrealdb_types::Error::connection(
287 "the server went away".to_string(),
288 surrealdb_types::ConnectionError::ConnectionFailed,
289 )));
290
291 let error = sessions.resolve(id).await.expect_err("the session was never established");
292 assert!(error.is_connection(), "the failure must still read as a connection failure");
293 }
294
295 #[tokio::test]
298 async fn closing_fails_the_waiting_and_everything_after() {
299 let sessions = registry();
300 let waiting = Uuid::new_v4();
301
302 let waiter = {
303 let sessions = Arc::clone(&sessions);
304 tokio::spawn(async move { sessions.resolve(waiting).await })
305 };
306 tokio::task::yield_now().await;
307 sessions.close();
308
309 assert!(matches!(waiter.await.unwrap(), Err(SessionError::NotFound(_))));
310 assert!(matches!(sessions.resolve(Uuid::new_v4()).await, Err(SessionError::NotFound(_))));
311 }
312}