#[cfg(not(target_family = "wasm"))]
pub(crate) mod native;
#[cfg(target_family = "wasm")]
pub(crate) mod wasm;
use std::marker::PhantomData;
use std::sync::atomic::{AtomicI64, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use async_channel::Sender;
use futures::{Sink, SinkExt};
use surrealdb_rpc::{DbResponse, DbResult, QueryResult, QueryResultBuilder, Token};
use surrealdb_types::{
AuthError, ConnectionError, Error as TypesError, NotAllowedError, SerializationError,
ValidationError,
};
use tokio::sync::RwLock;
use uuid::Uuid;
use crate::conn::{Command, RequestData, Route};
use crate::engine::remote::{RemoteCommand, RouterRequest};
use crate::engine::{SessionError, session_error_to_error};
use crate::opt::IntoEndpoint;
use crate::types::{Action, Array, HashMap, Notification, Number, SurrealValue, Value};
use crate::{Connect, Error, Surreal};
pub(crate) const PATH: &str = "rpc";
const PING_INTERVAL: Duration = Duration::from_secs(5);
#[derive(Debug, Clone)]
struct PendingRequest {
command: Option<Command>,
response_channel: Sender<Result<Vec<QueryResult>, TypesError>>,
}
struct SessionState {
pending_requests: HashMap<i64, PendingRequest>,
live_queries: HashMap<Uuid, Sender<crate::Result<Notification>>>,
replay: boxcar::Vec<Command>,
replay_cursor: Mutex<Option<ReplayCursor>>,
deferred: Mutex<Vec<Route>>,
last_id: AtomicI64,
}
#[derive(Debug, Clone, Copy)]
struct ReplayCursor {
awaiting: i64,
index: usize,
refreshed: bool,
}
impl Default for SessionState {
fn default() -> Self {
Self {
pending_requests: HashMap::new(),
live_queries: HashMap::new(),
replay: boxcar::Vec::new(),
replay_cursor: Mutex::new(None),
deferred: Mutex::new(Vec::new()),
last_id: AtomicI64::new(0),
}
}
}
impl Clone for SessionState {
fn clone(&self) -> Self {
Self {
replay: self.replay.clone(),
pending_requests: HashMap::new(),
live_queries: HashMap::new(),
replay_cursor: Mutex::new(None),
deferred: Mutex::new(Vec::new()),
last_id: AtomicI64::new(0),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum HandleResult {
Disconnected,
Ok,
}
trait WsMessage: Sized + Clone + Unpin + Send {
fn binary(payload: Vec<u8>) -> Self;
fn as_binary(&self) -> Option<&[u8]>;
fn should_process(&self) -> bool {
true
}
fn log_description(&self) -> &'static str {
"message"
}
}
fn serialize_request<M: WsMessage>(request: RouterRequest) -> M {
let request_value = request.into_value();
let payload = surrealdb_types::encode(&request_value).expect("router request should serialize");
M::binary(payload)
}
fn create_ping_message<M: WsMessage>() -> M {
let request = Command::Health
.into_router_request(None, None)
.expect("HEALTH command should convert to router request");
serialize_request(request)
}
fn create_kill_message<M: WsMessage>(live_query_id: Uuid, session_id: Uuid) -> M {
let request = Command::Kill {
uuid: live_query_id,
}
.into_router_request(None, Some(session_id))
.expect("KILL command should convert to router request");
serialize_request(request)
}
async fn send_message<M, S, E>(sink: &RwLock<S>, message: M) -> Result<(), E>
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
{
sink.write().await.send(message).await
}
async fn handle_route<M, S, E>(
route: Route,
max_message_size: Option<usize>,
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
sink: &RwLock<S>,
) -> HandleResult
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
let session_id = route.request.session_id;
let session_state = match sessions.get(&session_id) {
Some(Ok(state)) => state,
Some(Err(error)) => {
if route.response.send(Err(session_error_to_error(error))).await.is_err() {
trace!("Receiver dropped");
}
return HandleResult::Ok;
}
None => {
let error = session_error_to_error(SessionError::NotFound(session_id));
if route.response.send(Err(error)).await.is_err() {
trace!("Receiver dropped");
}
return HandleResult::Ok;
}
};
if replay_cursor(&session_state).is_some() {
park_route(&session_state, route);
return HandleResult::Ok;
}
dispatch_route::<M, S, E>(route, max_message_size, &session_state, sink).await
}
async fn dispatch_route<M, S, E>(
Route {
request,
response,
}: Route,
max_message_size: Option<usize>,
session_state: &Arc<SessionState>,
sink: &RwLock<S>,
) -> HandleResult
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
let RequestData {
command,
session_id,
} = request;
let id = session_state.last_id.fetch_add(1, Ordering::SeqCst);
if session_state.pending_requests.contains_key(&id) {
let error = Error::validation(
format!("Duplicate request ID: {id}"),
ValidationError::InvalidParams,
);
if response.send(Err(error)).await.is_err() {
trace!("Receiver dropped");
}
return HandleResult::Ok;
}
match command {
Command::SubscribeLive {
ref uuid,
ref notification_sender,
} => {
session_state.live_queries.insert(*uuid, notification_sender.clone());
if response.send(Ok(vec![QueryResultBuilder::instant_none()])).await.is_err() {
trace!("Receiver dropped");
}
return HandleResult::Ok;
}
Command::Kill {
ref uuid,
} => {
session_state.live_queries.remove(uuid);
}
_ => {}
}
let Some(router_request) = command.clone().into_router_request(Some(id), Some(session_id))
else {
response
.send(Err(Error::internal(
"The protocol or storage engine does not support backups on this architecture"
.to_string(),
)))
.await
.ok();
return HandleResult::Ok;
};
let message: M = serialize_request(router_request);
if let Some(max_size) = max_message_size
&& let Some(binary) = message.as_binary()
&& binary.len() > max_size
{
if response
.send(Err(Error::validation(
format!("Message too long: {}", binary.len()),
ValidationError::InvalidParams,
)))
.await
.is_err()
{
trace!("Receiver dropped");
}
return HandleResult::Ok;
}
match send_message(sink, message).await {
Ok(_) => {
session_state.pending_requests.insert(
id,
PendingRequest {
command: if command.replayable() {
Some(command)
} else {
None
},
response_channel: response,
},
);
}
Err(error) => {
let err = Error::connection(
format!("WebSocket error: {:?}", error),
ConnectionError::ConnectionFailed,
);
if response.send(Err(err)).await.is_err() {
trace!("Receiver dropped");
}
return HandleResult::Disconnected;
}
}
HandleResult::Ok
}
async fn handle_response<M, S, E>(
message: &M,
max_message_size: Option<usize>,
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
sink: &RwLock<S>,
) -> HandleResult
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
if !message.should_process() {
trace!("Received {}", message.log_description());
return HandleResult::Ok;
}
let Some(binary) = message.as_binary() else {
trace!("Received non-binary message");
return HandleResult::Ok;
};
match surrealdb_rpc::db_response_from_bytes(binary) {
Ok(response) => {
handle_db_response::<M, S, E>(response, max_message_size, sessions, sink).await
}
Err(error) => {
handle_parse_error(
Error::serialization(error.to_string(), SerializationError::Deserialization),
binary,
sessions,
)
.await
}
}
}
async fn handle_db_response<M, S, E>(
response: DbResponse,
max_message_size: Option<usize>,
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
sink: &RwLock<S>,
) -> HandleResult
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
let Some(session_id) = response.session_id else {
return HandleResult::Ok;
};
let session_state = match sessions.get(&session_id) {
Some(Ok(state)) => state,
_ => return HandleResult::Ok,
};
match response.id {
Some(id) => {
if let Value::Number(Number::Int(id_num)) = id {
handle_response_with_id::<M, S, E>(
id_num,
response.result,
session_id,
&session_state,
max_message_size,
sessions,
sink,
)
.await
} else {
HandleResult::Ok
}
}
None => {
handle_live_notification::<M, S, E>(response.result, session_id, &session_state, sink)
.await
}
}
}
async fn handle_response_with_id<M, S, E>(
id: i64,
result: Result<DbResult, TypesError>,
session_id: Uuid,
session_state: &Arc<SessionState>,
max_message_size: Option<usize>,
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
sink: &RwLock<S>,
) -> HandleResult
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
if let Some(cursor) = replay_cursor(session_state)
&& cursor.awaiting == id
{
if let Err(error) = result {
if !cursor.refreshed
&& let Some(Command::Authenticate {
token,
}) = session_state.replay.get(cursor.index)
&& let Token::WithRefresh {
..
} = token && error
.not_allowed_details()
.is_some_and(|a| matches!(a, NotAllowedError::Auth(AuthError::TokenExpired)))
{
let refresh_request = RouterRequest {
id: Some(id),
method: "authenticate",
params: Some(Value::Array(Array::from(vec![token.clone().into_value()]))),
txn: None,
session_id: Some(session_id),
};
let message: M = serialize_request(refresh_request);
if let Err(send_error) = send_message(sink, message).await {
trace!("failed to send refresh query to the server; {send_error:?}");
fail_replay(session_state, sessions, session_id, error).await;
} else {
set_replay_cursor(
session_state,
Some(ReplayCursor {
refreshed: true,
..cursor
}),
);
}
return HandleResult::Ok;
}
fail_replay(session_state, sessions, session_id, error).await;
return HandleResult::Ok;
}
if send_replay_command::<M, S, E>(session_id, session_state, cursor.index + 1, sink).await {
return HandleResult::Ok;
}
return flush_deferred_routes::<M, S, E>(session_state, max_message_size, sink).await;
}
let Some(mut pending) = session_state.pending_requests.take(&id) else {
warn!("got response for request with id '{id}', which was not in pending requests");
return HandleResult::Ok;
};
match result {
Ok(DbResult::Query(results)) => {
if let Some(command) = pending.command {
super::record_replayable(&session_state.replay, command);
}
if let Err(err) = pending.response_channel.send(Ok(results)).await {
tracing::error!("Failed to send query results to channel: {err:?}");
}
}
Ok(DbResult::Live(_)) => {
tracing::error!("Unexpected live query result in response");
}
Ok(DbResult::Other(mut value)) => {
if let Some(command) = pending.command {
if let Command::Authenticate {
token,
..
} = &command
{
value = token.clone().into_value();
}
super::record_replayable(&session_state.replay, command);
}
let result = QueryResultBuilder::started_now().finish_with_result(Ok(value));
if let Err(err) = pending.response_channel.send(Ok(vec![result])).await {
tracing::error!("Failed to send query results to channel: {err:?}");
}
}
Err(error) => {
if let Some(Command::Authenticate {
token,
..
}) = pending.command
&& let Token::WithRefresh {
..
} = &token && error
.not_allowed_details()
.is_some_and(|a| matches!(a, NotAllowedError::Auth(AuthError::TokenExpired)))
{
let refresh_request = RouterRequest {
id: Some(id),
method: "authenticate",
params: Some(Value::Array(Array::from(vec![token.into_value()]))),
txn: None,
session_id: Some(session_id),
};
let message: M = serialize_request(refresh_request);
match send_message(sink, message).await {
Err(send_error) => {
trace!("failed to send refresh query to the server; {send_error:?}");
pending.response_channel.send(Err(error)).await.ok();
}
Ok(..) => {
pending.command = None;
session_state.pending_requests.insert(id, pending);
}
}
return HandleResult::Ok;
}
pending.response_channel.send(Err(error)).await.ok();
}
}
HandleResult::Ok
}
fn park_route(session_state: &SessionState, route: Route) {
match session_state.deferred.lock() {
Ok(mut deferred) => deferred.push(route),
Err(poisoned) => poisoned.into_inner().push(route),
}
}
fn take_deferred_routes(session_state: &SessionState) -> Vec<Route> {
match session_state.deferred.lock() {
Ok(mut deferred) => std::mem::take(&mut *deferred),
Err(poisoned) => std::mem::take(&mut *poisoned.into_inner()),
}
}
async fn flush_deferred_routes<M, S, E>(
session_state: &Arc<SessionState>,
max_message_size: Option<usize>,
sink: &RwLock<S>,
) -> HandleResult
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
let mut outcome = HandleResult::Ok;
for route in take_deferred_routes(session_state) {
if let HandleResult::Disconnected =
dispatch_route::<M, S, E>(route, max_message_size, session_state, sink).await
{
outcome = HandleResult::Disconnected;
}
}
outcome
}
async fn fail_deferred_routes(session_state: &SessionState, error: TypesError) {
for route in take_deferred_routes(session_state) {
route.response.send(Err(error.clone())).await.ok();
}
}
async fn fail_replay(
session_state: &SessionState,
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
session_id: Uuid,
error: TypesError,
) {
set_replay_cursor(session_state, None);
sessions.insert(session_id, Err(SessionError::Remote(error.to_string())));
fail_deferred_routes(session_state, error).await;
}
async fn handle_live_notification<M, S, E>(
result: Result<DbResult, TypesError>,
session_id: Uuid,
session_state: &Arc<SessionState>,
sink: &RwLock<S>,
) -> HandleResult
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
if let Ok(DbResult::Live(notification)) = result {
let live_query_id = notification.id.into_inner();
let ended = matches!(notification.action, Action::Killed);
let registration = match ended {
true => session_state.live_queries.take(&live_query_id),
false => session_state.live_queries.get(&live_query_id),
};
if let Some(sender) = registration
&& sender.send(Ok(notification)).await.is_err()
&& !ended
{
session_state.live_queries.remove(&live_query_id);
let kill: M = create_kill_message(live_query_id, session_id);
if let Err(error) = send_message(sink, kill).await {
trace!("failed to send kill query to the server; {error:?}");
return HandleResult::Disconnected;
}
}
}
HandleResult::Ok
}
async fn handle_parse_error(
error: crate::Error,
binary: &[u8],
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
) -> HandleResult {
#[derive(SurrealValue)]
#[surreal(crate = "crate::types")]
struct ErrorResponse {
id: Option<Value>,
#[surreal(rename = "session")]
session_id: Option<Uuid>,
}
match surrealdb_types::decode::<ErrorResponse>(binary) {
Ok(ErrorResponse {
id,
session_id,
}) => {
let Some(session_id) = session_id else {
return HandleResult::Ok;
};
let session_state = match sessions.get(&session_id) {
Some(Ok(state)) => state,
_ => return HandleResult::Ok,
};
match id {
Some(Value::Number(Number::Int(id_num))) => {
if let Some(pending) = session_state.pending_requests.take(&id_num) {
let _ = pending.response_channel.send(Err(error)).await;
} else if replay_cursor(&session_state).is_some_and(|c| c.awaiting == id_num) {
fail_replay(&session_state, sessions, session_id, error).await;
} else {
warn!(
"got response for request with id '{id_num}', which was not in pending requests"
);
}
}
_ => {
if replay_cursor(&session_state).is_some() {
fail_replay(&session_state, sessions, session_id, error).await;
}
}
}
}
_ => {
error!("Failed to deserialise message, failing pending requests; {error:?}");
fail_all_pending_requests(sessions, error).await;
}
}
HandleResult::Ok
}
async fn fail_all_pending_requests(
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
error: crate::Error,
) {
for (session_id, session) in sessions.to_vec() {
let Ok(session_state) = session else {
continue;
};
for (id, _) in session_state.pending_requests.to_vec() {
if let Some(pending) = session_state.pending_requests.take(&id) {
pending.response_channel.send(Err(error.clone())).await.ok();
}
}
if replay_cursor(&session_state).is_some() {
fail_replay(&session_state, sessions, session_id, error.clone()).await;
}
}
}
async fn replay_session<M, S, E>(
session_id: Uuid,
session_state: &SessionState,
sink: &RwLock<S>,
) -> crate::Result<()>
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
send_replay_command::<M, S, E>(session_id, session_state, 0, sink).await;
Ok(())
}
async fn send_replay_command<M, S, E>(
session_id: Uuid,
session_state: &SessionState,
index: usize,
sink: &RwLock<S>,
) -> bool
where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
let Some(command) = session_state.replay.get(index) else {
set_replay_cursor(session_state, None);
return false;
};
let id = session_state.last_id.fetch_add(1, Ordering::SeqCst);
set_replay_cursor(
session_state,
Some(ReplayCursor {
awaiting: id,
index,
refreshed: false,
}),
);
let request = command
.clone()
.into_router_request(Some(id), Some(session_id))
.expect("replay commands should always convert to route requests");
let message: M = serialize_request(request);
if let Err(error) = send_message(sink, message).await {
set_replay_cursor(session_state, None);
debug!("{:?}", error);
return false;
}
true
}
fn replay_cursor(session_state: &SessionState) -> Option<ReplayCursor> {
match session_state.replay_cursor.lock() {
Ok(cursor) => *cursor,
Err(poisoned) => *poisoned.into_inner(),
}
}
fn set_replay_cursor(session_state: &SessionState, cursor: Option<ReplayCursor>) {
match session_state.replay_cursor.lock() {
Ok(mut slot) => *slot = cursor,
Err(poisoned) => *poisoned.into_inner() = cursor,
}
}
async fn handle_session_initial<M, S, E>(
session_id: Uuid,
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
sink: &RwLock<S>,
) where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
let session_state = Arc::new(SessionState::default());
session_state.replay.push(Command::Attach {
session_id,
});
sessions.insert(session_id, Ok(Arc::clone(&session_state)));
if let Err(error) = replay_session::<M, S, E>(session_id, &session_state, sink).await {
sessions.insert(session_id, Err(SessionError::Remote(error.to_string())));
}
}
async fn handle_session_clone<M, S, E>(
old: Uuid,
new: Uuid,
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
sink: &RwLock<S>,
) where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
match sessions.get(&old) {
Some(Ok(session_state)) => {
let mut session_state = session_state.as_ref().clone();
if let Some(cmd) = session_state.replay.get_mut(0) {
*cmd = Command::Attach {
session_id: new,
};
}
let session_state = Arc::new(session_state);
sessions.insert(new, Ok(Arc::clone(&session_state)));
if let Err(error) = replay_session::<M, S, E>(new, &session_state, sink).await {
sessions.insert(new, Err(SessionError::Remote(error.to_string())));
}
}
Some(Err(error)) => {
sessions.insert(new, Err(error));
}
None => {
sessions.insert(new, Err(SessionError::NotFound(old)));
}
}
}
async fn handle_session_drop<M, S, E>(
session_id: Uuid,
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
sink: &RwLock<S>,
) where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
if sessions.get(&session_id).is_some() {
let request = Command::Detach {
session_id,
}
.into_router_request(None, Some(session_id))
.expect("detach should always convert to a route request");
let message: M = serialize_request(request);
if let Err(error) = send_message(sink, message).await {
debug!("{:?}", error);
}
}
sessions.remove(&session_id);
}
async fn handle_session<M, S, E>(
session_id: crate::SessionId,
sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>,
sink: &RwLock<S>,
) where
M: WsMessage,
S: Sink<M, Error = E> + Unpin,
E: std::fmt::Debug,
{
match session_id {
crate::SessionId::Initial(id) => {
handle_session_initial::<M, S, E>(id, sessions, sink).await
}
crate::SessionId::Clone {
old,
new,
} => handle_session_clone::<M, S, E>(old, new, sessions, sink).await,
crate::SessionId::Drop(id) => handle_session_drop::<M, S, E>(id, sessions, sink).await,
}
}
async fn clear_pending_requests(sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>) {
for state in sessions.values().into_iter().flatten() {
for request in state.pending_requests.values() {
let err = crate::Error::connection(
"Connection reset".to_string(),
surrealdb_types::ConnectionError::ConnectionFailed,
);
request.response_channel.send(Err(err)).await.ok();
request.response_channel.close();
}
state.pending_requests.clear();
let err = crate::Error::connection(
"Connection reset".to_string(),
surrealdb_types::ConnectionError::ConnectionFailed,
);
fail_deferred_routes(&state, err).await;
}
}
async fn clear_live_queries(sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>) {
for state in sessions.values().into_iter().flatten() {
for sender in state.live_queries.values() {
let err = crate::Error::connection(
"Connection reset".to_string(),
surrealdb_types::ConnectionError::ConnectionFailed,
);
sender.send(Err(err)).await.ok();
sender.close();
}
state.live_queries.clear();
}
}
async fn reset_sessions(sessions: &HashMap<Uuid, Result<Arc<SessionState>, SessionError>>) {
tokio::join!(clear_pending_requests(sessions), clear_live_queries(sessions));
}
#[derive(Debug)]
pub struct Ws;
#[derive(Debug)]
pub struct Wss;
#[derive(Debug, Clone)]
pub struct Client(());
impl Surreal<Client> {
pub fn connect<P>(
&self,
address: impl IntoEndpoint<P, Client = Client>,
) -> Connect<Client, ()> {
Connect {
surreal: Arc::clone(&self.inner).into(),
address: address.into_endpoint(),
capacity: 0,
response_type: PhantomData,
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use surrealdb_rpc::{DbResult, QueryResult, Token};
use surrealdb_types::{AuthError, Error as TypesError, NotAllowedError};
use tokio::sync::RwLock;
use uuid::Uuid;
use super::{
HandleResult, PendingRequest, SessionState, WsMessage, fail_all_pending_requests,
handle_live_notification, handle_parse_error, handle_response_with_id, handle_route,
replay_cursor, replay_session,
};
use crate::conn::{Command, RequestData, Route};
use crate::engine::SessionError;
use crate::types::{Action, HashMap, Notification, Number, Value};
type Sessions = HashMap<Uuid, Result<Arc<SessionState>, SessionError>>;
fn no_sessions() -> Sessions {
HashMap::new()
}
#[derive(Clone)]
struct MockMessage;
impl WsMessage for MockMessage {
fn binary(_payload: Vec<u8>) -> Self {
MockMessage
}
fn as_binary(&self) -> Option<&[u8]> {
None
}
}
#[tokio::test]
async fn handle_response_removes_pending_request() {
let session_state = Arc::new(SessionState::default());
let sessions = no_sessions();
let session_id = Uuid::new_v4();
let request_id: i64 = 1;
let (sender, receiver) = async_channel::bounded(1);
session_state.pending_requests.insert(
request_id,
PendingRequest {
command: None,
response_channel: sender,
},
);
assert_eq!(session_state.pending_requests.len(), 1);
let sink = RwLock::new(futures::sink::drain::<MockMessage>());
let result = handle_response_with_id::<MockMessage, _, _>(
request_id,
Ok(DbResult::Other(Value::None)),
session_id,
&session_state,
None,
&sessions,
&sink,
)
.await;
assert_eq!(result, HandleResult::Ok);
assert!(
session_state.pending_requests.is_empty(),
"pending request should be removed after handling response"
);
let response = receiver.recv().await.unwrap();
assert!(response.is_ok());
}
#[tokio::test]
async fn handle_response_error_removes_pending_request() {
let session_state = Arc::new(SessionState::default());
let sessions = no_sessions();
let session_id = Uuid::new_v4();
let request_id: i64 = 1;
let (sender, receiver) = async_channel::bounded(1);
session_state.pending_requests.insert(
request_id,
PendingRequest {
command: None,
response_channel: sender,
},
);
assert_eq!(session_state.pending_requests.len(), 1);
let sink = RwLock::new(futures::sink::drain::<MockMessage>());
let error = TypesError::internal("test error".to_string());
let result = handle_response_with_id::<MockMessage, _, _>(
request_id,
Err(error),
session_id,
&session_state,
None,
&sessions,
&sink,
)
.await;
assert_eq!(result, HandleResult::Ok);
assert!(
session_state.pending_requests.is_empty(),
"pending request should be removed after handling error response"
);
let response = receiver.recv().await.unwrap();
assert!(response.is_err());
}
#[tokio::test]
async fn handle_multiple_responses_cleans_up_all_entries() {
let session_state = Arc::new(SessionState::default());
let sessions = no_sessions();
let session_id = Uuid::new_v4();
let sink = RwLock::new(futures::sink::drain::<MockMessage>());
let mut receivers = Vec::new();
for id in 0..100i64 {
let (sender, receiver) = async_channel::bounded(1);
session_state.pending_requests.insert(
id,
PendingRequest {
command: None,
response_channel: sender,
},
);
receivers.push(receiver);
}
assert_eq!(session_state.pending_requests.len(), 100);
for id in 0..100i64 {
handle_response_with_id::<MockMessage, _, _>(
id,
Ok(DbResult::Other(Value::None)),
session_id,
&session_state,
None,
&sessions,
&sink,
)
.await;
}
assert!(
session_state.pending_requests.is_empty(),
"all pending requests should be removed, but {} remain",
session_state.pending_requests.len()
);
for receiver in &receivers {
let response = receiver.recv().await.unwrap();
assert!(response.is_ok());
}
}
fn session_with_replay(commands: Vec<Command>) -> (Uuid, Arc<SessionState>, Sessions) {
let session_id = Uuid::new_v4();
let session_state = Arc::new(SessionState::default());
for command in commands {
session_state.replay.push(command);
}
let sessions = HashMap::new();
sessions.insert(session_id, Ok(Arc::clone(&session_state)));
(session_id, session_state, sessions)
}
fn set_cmd(key: &str, value: i64) -> Command {
Command::Set {
key: key.to_string(),
value: Value::Number(Number::Int(value)),
}
}
async fn ack_replay<S>(
session_id: Uuid,
session_state: &Arc<SessionState>,
sessions: &Sessions,
sink: &RwLock<S>,
result: Result<DbResult, TypesError>,
) where
S: futures::Sink<MockMessage, Error = std::convert::Infallible> + Unpin,
{
let cursor = replay_cursor(session_state).expect("a replay command should be in flight");
handle_response_with_id::<MockMessage, _, _>(
cursor.awaiting,
result,
session_id,
session_state,
None,
sessions,
sink,
)
.await;
}
fn route_for(
session_id: Uuid,
) -> (Route, async_channel::Receiver<Result<Vec<QueryResult>, TypesError>>) {
let (response, receiver) = async_channel::bounded(1);
let route = Route {
request: RequestData {
command: Command::Health,
session_id,
},
response,
};
(route, receiver)
}
#[tokio::test]
async fn route_waits_for_replay_acknowledgement() {
let (session_id, session_state, sessions) = session_with_replay(vec![set_cmd("x", 1)]);
let sink = RwLock::new(Vec::<MockMessage>::new());
replay_session::<MockMessage, _, _>(session_id, &session_state, &sink).await.unwrap();
assert_eq!(sink.read().await.len(), 1, "the replay command should be on the wire");
let (route, _receiver) = route_for(session_id);
assert_eq!(
handle_route::<MockMessage, _, _>(route, None, &sessions, &sink).await,
HandleResult::Ok
);
assert_eq!(
sink.read().await.len(),
1,
"request must not reach the wire before the session's setup is acknowledged"
);
assert_eq!(session_state.deferred.lock().unwrap().len(), 1, "request should be parked");
ack_replay(session_id, &session_state, &sessions, &sink, Ok(DbResult::Other(Value::None)))
.await;
assert!(replay_cursor(&session_state).is_none(), "replay should be finished");
assert!(
session_state.deferred.lock().unwrap().is_empty(),
"parked request should be released once setup is acknowledged"
);
assert_eq!(sink.read().await.len(), 2, "released request should be sent");
assert_eq!(session_state.pending_requests.len(), 1, "released request should be pending");
}
#[tokio::test]
async fn replay_commands_are_sent_one_at_a_time() {
let (session_id, session_state, sessions) =
session_with_replay(vec![set_cmd("x", 1), set_cmd("x", 2)]);
let sink = RwLock::new(Vec::<MockMessage>::new());
replay_session::<MockMessage, _, _>(session_id, &session_state, &sink).await.unwrap();
assert_eq!(
sink.read().await.len(),
1,
"only the first command may be in flight; the second would be free to overtake it"
);
let first = replay_cursor(&session_state).expect("first command in flight");
ack_replay(session_id, &session_state, &sessions, &sink, Ok(DbResult::Other(Value::None)))
.await;
assert_eq!(
sink.read().await.len(),
2,
"the second command should follow its predecessor's acknowledgement"
);
let second = replay_cursor(&session_state).expect("second command in flight");
assert_ne!(second.awaiting, first.awaiting, "each command needs its own request id");
ack_replay(session_id, &session_state, &sessions, &sink, Ok(DbResult::Other(Value::None)))
.await;
assert!(
replay_cursor(&session_state).is_none(),
"the session should be ready once the log is exhausted"
);
}
fn refreshable_auth_cmd() -> Command {
Command::Authenticate {
token: Token::WithRefresh {
access: "expired-access".to_string(),
refresh: "valid-refresh".to_string(),
},
}
}
fn token_expired() -> TypesError {
TypesError::not_allowed(
"token expired".to_string(),
NotAllowedError::Auth(AuthError::TokenExpired),
)
}
#[tokio::test]
async fn expired_token_during_replay_is_refreshed_not_poisoned() {
let (session_id, session_state, sessions) =
session_with_replay(vec![refreshable_auth_cmd(), set_cmd("x", 1)]);
let sink = RwLock::new(Vec::<MockMessage>::new());
replay_session::<MockMessage, _, _>(session_id, &session_state, &sink).await.unwrap();
let first = replay_cursor(&session_state).expect("authenticate in flight");
ack_replay(session_id, &session_state, &sessions, &sink, Err(token_expired())).await;
assert!(
matches!(sessions.get(&session_id), Some(Ok(_))),
"a refreshable expiry must not poison the session"
);
let retry = replay_cursor(&session_state).expect("replay should still be in flight");
assert_eq!(retry.index, first.index, "the retry re-sends the same command");
assert!(retry.refreshed, "the retry should be recorded so it happens only once");
assert_eq!(sink.read().await.len(), 2, "the refreshed authenticate should be sent");
ack_replay(session_id, &session_state, &sessions, &sink, Ok(DbResult::Other(Value::None)))
.await;
let next = replay_cursor(&session_state).expect("set should now be in flight");
assert_eq!(next.index, first.index + 1, "the replay should advance past authenticate");
assert!(!next.refreshed, "a fresh command starts with no retry spent");
}
#[tokio::test]
async fn expired_token_is_refreshed_only_once() {
let (session_id, session_state, sessions) =
session_with_replay(vec![refreshable_auth_cmd()]);
let sink = RwLock::new(Vec::<MockMessage>::new());
replay_session::<MockMessage, _, _>(session_id, &session_state, &sink).await.unwrap();
ack_replay(session_id, &session_state, &sessions, &sink, Err(token_expired())).await;
assert!(replay_cursor(&session_state).is_some_and(|c| c.refreshed));
ack_replay(session_id, &session_state, &sessions, &sink, Err(token_expired())).await;
assert!(replay_cursor(&session_state).is_none(), "the replay should be abandoned");
assert!(
matches!(sessions.get(&session_id), Some(Err(_))),
"a refresh that cannot recover should poison the session"
);
assert_eq!(sink.read().await.len(), 2, "no further refresh attempts");
}
#[tokio::test]
async fn response_without_usable_id_releases_a_parked_replay() {
let (session_id, session_state, sessions) = session_with_replay(vec![set_cmd("x", 1)]);
let sink = RwLock::new(Vec::<MockMessage>::new());
replay_session::<MockMessage, _, _>(session_id, &session_state, &sink).await.unwrap();
let (route, receiver) = route_for(session_id);
handle_route::<MockMessage, _, _>(route, None, &sessions, &sink).await;
assert_eq!(session_state.deferred.lock().unwrap().len(), 1);
let mut envelope = crate::types::Object::new();
envelope.insert("id".to_string(), Value::None);
envelope.insert("session".to_string(), Value::Uuid(session_id.into()));
let binary = surrealdb_types::encode(&Value::Object(envelope)).unwrap();
handle_parse_error(TypesError::internal("unreadable".to_string()), &binary, &sessions)
.await;
assert!(replay_cursor(&session_state).is_none(), "the stuck replay should be abandoned");
assert!(
receiver.recv().await.unwrap().is_err(),
"the parked request should be failed, not left hanging"
);
}
#[tokio::test]
async fn undecodable_response_releases_a_parked_replay() {
let (session_id, session_state, sessions) = session_with_replay(vec![set_cmd("x", 1)]);
let sink = RwLock::new(Vec::<MockMessage>::new());
replay_session::<MockMessage, _, _>(session_id, &session_state, &sink).await.unwrap();
let (route, receiver) = route_for(session_id);
handle_route::<MockMessage, _, _>(route, None, &sessions, &sink).await;
assert_eq!(session_state.deferred.lock().unwrap().len(), 1);
fail_all_pending_requests(&sessions, TypesError::internal("undecodable".to_string())).await;
assert!(replay_cursor(&session_state).is_none(), "the stuck replay should be abandoned");
assert!(
receiver.recv().await.unwrap().is_err(),
"the parked request should be failed, not left hanging"
);
assert!(
matches!(sessions.get(&session_id), Some(Err(_))),
"a session whose setup cannot be confirmed should be poisoned"
);
}
#[tokio::test]
async fn failed_replay_acknowledgement_fails_parked_requests() {
let (session_id, session_state, sessions) = session_with_replay(vec![set_cmd("x", 1)]);
let sink = RwLock::new(Vec::<MockMessage>::new());
replay_session::<MockMessage, _, _>(session_id, &session_state, &sink).await.unwrap();
let sent_during_replay = sink.read().await.len();
let (route, receiver) = route_for(session_id);
handle_route::<MockMessage, _, _>(route, None, &sessions, &sink).await;
assert_eq!(session_state.deferred.lock().unwrap().len(), 1);
ack_replay(
session_id,
&session_state,
&sessions,
&sink,
Err(TypesError::internal("setup rejected".to_string())),
)
.await;
assert_eq!(
sink.read().await.len(),
sent_during_replay,
"a rejected setup must not release the request"
);
assert!(
receiver.recv().await.unwrap().is_err(),
"the parked request should be failed, not left hanging"
);
assert!(
matches!(sessions.get(&session_id), Some(Err(_))),
"the session should be poisoned once its setup is rejected"
);
}
async fn deliver(action: Action) -> (bool, bool) {
let session_state = Arc::new(SessionState::default());
let live = Uuid::new_v4();
let (sender, subscriber) = async_channel::unbounded();
session_state.live_queries.insert(live, sender);
let sink = RwLock::new(futures::sink::drain::<MockMessage>());
let result = handle_live_notification::<MockMessage, _, _>(
Ok(DbResult::Live(Notification::new(
live.into(),
None,
action,
Value::None,
Value::None,
))),
Uuid::new_v4(),
&session_state,
&sink,
)
.await;
assert_eq!(result, HandleResult::Ok);
(subscriber.try_recv().is_ok(), session_state.live_queries.contains_key(&live))
}
#[tokio::test]
async fn a_killed_notification_takes_the_registration_with_it() {
let (delivered, registered) = deliver(Action::Killed).await;
assert!(delivered, "the subscriber is still owed the end of its subscription");
assert!(!registered, "the registration outlived the subscription it names");
}
#[tokio::test]
async fn a_change_notification_leaves_the_registration_in_place() {
let (delivered, registered) = deliver(Action::Create).await;
assert!(delivered, "the change reaches the subscriber");
assert!(registered, "the subscription is still running and still the session's");
}
}