use core::fmt;
use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::close_code::AGAIN;
use axum::extract::ws::{CloseFrame, Message, WebSocket};
use bytes::Bytes;
use dashmap::DashMap;
use futures::stream::FuturesUnordered;
use futures::{Sink, SinkExt, StreamExt};
use opentelemetry::Context as TelemetryContext;
use opentelemetry::trace::FutureExt;
use surrealdb_core::dbs::Session;
use surrealdb_core::kvs::{Datastore, LockType, Transaction, TransactionType};
use surrealdb_core::mem::ALLOC;
use surrealdb_core::rpc::format::Format;
use surrealdb_core::rpc::{DbResponse, DbResult, Method, RpcProtocol};
use surrealdb_types::{Array, Error as TypesError, HashMap, Value};
use tokio::sync::RwLock;
use tokio::sync::mpsc::{Receiver, Sender, channel};
use tokio::task::JoinSet;
use tokio_util::sync::CancellationToken;
use tracing::{Instrument, Span};
use uuid::Uuid;
use super::RpcState;
use crate::cnf::{
PKG_NAME, PKG_VERSION, WEBSOCKET_PING_FREQUENCY, WEBSOCKET_RESPONSE_BUFFER_SIZE,
WEBSOCKET_RESPONSE_CHANNEL_SIZE, WEBSOCKET_RESPONSE_FLUSH_PERIOD,
};
use crate::rpc::CONN_CLOSED_ERR;
use crate::rpc::format::WsFormat;
use crate::telemetry;
use crate::telemetry::metrics::ws::RequestContext;
use crate::telemetry::traces::rpc::span_for_request;
const SERVER_OVERLOADED: &str = "The server is unable to handle the request";
const SERVER_SHUTTING_DOWN: &str = "The server is gracefully shutting down";
pub struct Websocket {
pub(crate) id: Uuid,
pub(crate) format: Format,
pub(crate) state: Arc<RpcState>,
pub(crate) datastore: Arc<Datastore>,
pub(crate) sessions: HashMap<Option<Uuid>, Arc<RwLock<Session>>>,
pub(crate) transactions: DashMap<Uuid, Arc<Transaction>>,
pub(crate) shutdown: CancellationToken,
pub(crate) canceller: CancellationToken,
pub(crate) channel: Sender<Message>,
}
impl Websocket {
pub async fn serve(
id: Uuid,
ws: WebSocket,
format: Format,
session: Session,
datastore: Arc<Datastore>,
state: Arc<RpcState>,
) {
trace!("WebSocket {id} connected");
let (sender, receiver) = channel(*WEBSOCKET_RESPONSE_CHANNEL_SIZE);
let rpc = Arc::new(Websocket {
id,
format,
state: state.clone(),
shutdown: CancellationToken::new(),
canceller: CancellationToken::new(),
sessions: HashMap::new(),
transactions: DashMap::new(),
channel: sender.clone(),
datastore,
});
let session = session.with_rt(true);
rpc.set_session(None, Arc::new(RwLock::new(session)));
state.web_sockets.write().await.insert(id, rpc.clone());
telemetry::metrics::ws::on_connect();
let mut tasks = JoinSet::new();
match *WEBSOCKET_RESPONSE_BUFFER_SIZE > 0 {
true => {
let buffer = ws.buffer(*WEBSOCKET_RESPONSE_BUFFER_SIZE);
let (ws_sender, ws_receiver) = buffer.split();
tasks.spawn(Self::ping(rpc.clone(), sender.clone()));
tasks.spawn(Self::read(rpc.clone(), ws_receiver, sender.clone()));
tasks.spawn(Self::write(rpc.clone(), ws_sender, receiver));
}
false => {
let (ws_sender, ws_receiver) = ws.split();
tasks.spawn(Self::ping(rpc.clone(), sender.clone()));
tasks.spawn(Self::read(rpc.clone(), ws_receiver, sender.clone()));
tasks.spawn(Self::write(rpc.clone(), ws_sender, receiver));
}
}
while let Some(res) = tasks.join_next().await {
if let Err(err) = res {
error!("Error handling RPC connection: {err}");
}
}
std::mem::drop(sender);
trace!("WebSocket {id} disconnected");
rpc.cleanup_all_lqs().await;
state.web_sockets.write().await.remove(&id);
telemetry::metrics::ws::on_disconnect();
}
async fn ping(rpc: Arc<Websocket>, internal_sender: Sender<Message>) {
let mut interval = tokio::time::interval(WEBSOCKET_PING_FREQUENCY);
let canceller = rpc.canceller.clone();
loop {
tokio::select! {
biased;
_ = canceller.cancelled() => break,
_ = interval.tick() => {
let msg = Message::Ping(Bytes::from_static(b""));
if let Err(err) = internal_sender.send(msg).await {
if err.to_string() != CONN_CLOSED_ERR {
trace!("WebSocket error: {err}");
}
canceller.cancel();
break;
}
},
}
}
}
async fn write<S: SinkExt<Message> + Unpin>(
rpc: Arc<Websocket>,
mut socket: S,
mut internal_receiver: Receiver<Message>,
) where
<S as Sink<Message>>::Error: fmt::Display,
{
let canceller = rpc.canceller.clone();
let buffer = *WEBSOCKET_RESPONSE_BUFFER_SIZE > 0;
let period = Duration::from_millis(*WEBSOCKET_RESPONSE_FLUSH_PERIOD);
loop {
tokio::select! {
biased;
_ = canceller.cancelled() => break,
Some(res) = internal_receiver.recv() => {
let res = match buffer {
true => socket.feed(res).await,
false => socket.send(res).await
};
if let Err(err) = res {
if err.to_string() != CONN_CLOSED_ERR {
trace!("WebSocket error: {err}");
}
canceller.cancel();
break;
}
},
_ = tokio::time::sleep(period), if buffer => {
if let Err(err) = socket.flush().await {
if err.to_string() != CONN_CLOSED_ERR {
trace!("WebSocket error: {err}");
}
canceller.cancel();
break;
}
}
}
}
}
async fn read(
rpc: Arc<Websocket>,
mut socket: impl StreamExt<Item = Result<Message, axum::Error>> + Unpin,
internal_sender: Sender<Message>,
) {
let shutdown = rpc.shutdown.clone();
let canceller = rpc.canceller.clone();
let mut tasks = FuturesUnordered::new();
loop {
tokio::select! {
biased;
_ = tasks.next(), if !tasks.is_empty() => {},
_ = shutdown.cancelled() => break,
_ = canceller.cancelled() => break,
Some(msg) = socket.next() => match msg {
Ok(msg) => match msg {
Message::Text(_) | Message::Binary(_) => {
let chn = internal_sender.clone();
if ALLOC.is_beyond_threshold() {
Self::close_socket(rpc.clone(), chn).await;
break;
}
tasks.push(Self::handle_message(&rpc, msg, chn));
}
Message::Close(_) => {
if let Err(err) = internal_sender.send(Message::Close(None)).await {
trace!("WebSocket error when replying to the close message: {err}");
};
canceller.cancel();
break;
}
Message::Ping(_) => {
}
Message::Pong(_) => {
}
},
Err(err) => {
trace!("WebSocket error: {err}");
canceller.cancel();
break;
}
}
}
}
tokio::select! {
biased;
_ = canceller.cancelled() => (),
_ = shutdown.cancelled() => {
while tasks.next().await.is_some() {
}
},
}
canceller.cancel();
std::mem::drop(tasks);
}
async fn handle_message(rpc: &Arc<Websocket>, msg: Message, chn: Sender<Message>) {
let shutdown = rpc.shutdown.clone();
let canceller = rpc.canceller.clone();
let len = match msg {
Message::Text(ref msg) => msg.len(),
Message::Binary(ref msg) => msg.len(),
_ => 0,
};
let span = span_for_request(&rpc.id);
async move {
let span = Span::current();
let req_cx = RequestContext::default();
let otel_cx = Arc::new(TelemetryContext::new().with_value(req_cx.clone()));
match rpc.format.req_ws(msg) {
Ok(req) => {
span.record("rpc.method", req.method.to_str());
span.record("otel.name", format!("surrealdb.rpc/{}", req.method));
span.record(
"rpc.request_id",
req.id.clone().map(|id| format!("{id:?}")).unwrap_or_default(),
);
let otel_cx = Arc::new(TelemetryContext::current_with_value(
req_cx.with_method(req.method.to_str()).with_size(len),
));
tokio::select! {
biased;
_ = canceller.cancelled() => (),
_ = async move {
if shutdown.is_cancelled() {
crate::rpc::response::send(
DbResponse::failure(req.id, req.session_id.map(Into::into), TypesError::internal(SERVER_SHUTTING_DOWN.to_string())),
otel_cx.clone(),
rpc.format,
chn
)
.with_context(otel_cx.as_ref().clone())
.await;
}
else if ALLOC.is_beyond_threshold() {
crate::rpc::response::send(
DbResponse::failure(req.id, req.session_id.map(Into::into), TypesError::internal(SERVER_OVERLOADED.to_string())),
otel_cx.clone(),
rpc.format,
chn
)
.with_context(otel_cx.as_ref().clone())
.await;
}
else {
let result = Self::process_message(
rpc.clone(),
req.session_id.map(Into::into),
req.txn.map(Into::into),
req.method,
req.params,
)
.await;
crate::rpc::response::send(
match result {
Ok(result) => DbResponse::success(req.id, req.session_id.map(Into::into), result),
Err(err) => DbResponse::failure(req.id, req.session_id.map(Into::into), err),
},
otel_cx.clone(),
rpc.format,
chn
)
.with_context(otel_cx.as_ref().clone())
.await;
}
} => (),
}
}
Err(err) => {
crate::rpc::response::send(
DbResponse::failure(None, None, err),
otel_cx.clone(),
rpc.format,
chn
)
.with_context(otel_cx.as_ref().clone())
.await;
}
}
}
.instrument(span)
.await;
}
async fn process_message(
rpc: Arc<Websocket>,
session_id: Option<Uuid>,
txn: Option<Uuid>,
method: Method,
params: Array,
) -> Result<DbResult, TypesError> {
debug!("Process RPC request");
if !method.is_valid() {
return Err(TypesError::not_found(
"Method not found".to_string(),
Some(surrealdb_types::NotFoundError::Method {
name: method.to_string(),
}),
));
}
RpcProtocol::execute(rpc.as_ref(), txn, session_id, method, params).await
}
async fn close_socket(rpc: Arc<Websocket>, chn: Sender<Message>) {
warn!("The server is overloaded and is unable to process a WebSocket request");
let frame = CloseFrame {
code: AGAIN,
reason: SERVER_OVERLOADED.into(),
};
if let Err(err) = chn.send(Message::Close(Some(frame))).await {
debug!("WebSocket error when sending close message: {err}");
};
rpc.canceller.cancel();
}
}
impl RpcProtocol for Websocket {
fn kvs(&self) -> &Datastore {
&self.datastore
}
fn version_data(&self) -> DbResult {
let value = Value::String(format!("{PKG_NAME}-{}", *PKG_VERSION));
DbResult::Other(value)
}
fn session_map(&self) -> &HashMap<Option<Uuid>, Arc<RwLock<Session>>> {
&self.sessions
}
async fn get_tx(
&self,
id: Uuid,
) -> Result<Arc<surrealdb_core::kvs::Transaction>, surrealdb_types::Error> {
debug!("WebSocket get_tx called for transaction {id}");
self.transactions
.get(&id)
.map(|tx| {
debug!("Transaction {id} found in WebSocket transactions map");
tx.clone()
})
.ok_or_else(|| {
warn!(
"Transaction {id} not found in WebSocket transactions map (have {} transactions)",
self.transactions.len()
);
surrealdb_core::rpc::invalid_params("Transaction not found")
})
}
async fn set_tx(
&self,
id: Uuid,
tx: Arc<surrealdb_core::kvs::Transaction>,
) -> Result<(), surrealdb_types::Error> {
self.transactions.insert(id, tx);
Ok(())
}
const LQ_SUPPORT: bool = true;
async fn handle_live(&self, lqid: &Uuid, session_id: Option<Uuid>) {
self.state.live_queries.write().await.insert(*lqid, (self.id, session_id));
trace!("Registered live query {lqid} on websocket {}", self.id);
}
async fn handle_kill(&self, lqid: &Uuid) {
if let Some((id, session_id)) = self.state.live_queries.write().await.remove(lqid) {
if let Some(session_id) = session_id {
trace!("Unregistered live query {lqid} on websocket {id} for session {session_id}");
} else {
trace!("Unregistered live query {lqid} on websocket {id} for default session");
}
}
}
async fn cleanup_lqs(&self, session_id: Option<&Uuid>) {
let mut gc = Vec::new();
self.state.live_queries.write().await.retain(|key, value| {
if value.0 == self.id && value.1.as_ref() == session_id {
trace!("Removing live query: {key}");
gc.push(*key);
return false;
}
true
});
if let Err(err) = self.kvs().delete_queries(gc).await {
error!("Error handling RPC connection: {err}");
}
}
async fn cleanup_all_lqs(&self) {
let mut gc = Vec::new();
self.state.live_queries.write().await.retain(|key, value| {
if value.0 == self.id {
trace!("Removing live query: {key}");
gc.push(*key);
return false;
}
true
});
if let Err(err) = self.kvs().delete_queries(gc).await {
error!("Error handling RPC connection: {err}");
}
}
async fn begin(
&self,
_txn: Option<Uuid>,
_session_id: Option<Uuid>,
) -> Result<DbResult, surrealdb_types::Error> {
let tx = self
.kvs()
.transaction(TransactionType::Write, LockType::Optimistic)
.await
.map_err(surrealdb_core::rpc::types_error_from_anyhow)?;
let id = Uuid::now_v7();
debug!("WebSocket begin: created transaction {id}");
self.transactions.insert(id, Arc::new(tx));
debug!(
"WebSocket begin: stored transaction {id}, map now has {} transactions",
self.transactions.len()
);
Ok(DbResult::Other(Value::Uuid(surrealdb::types::Uuid::from(id))))
}
async fn commit(
&self,
_txn: Option<Uuid>,
_session_id: Option<Uuid>,
params: Array,
) -> Result<DbResult, surrealdb_types::Error> {
let mut params_vec = params.into_vec();
let Some(Value::Uuid(txn_id)) = params_vec.pop() else {
return Err(surrealdb_core::rpc::invalid_params("Expected transaction UUID"));
};
let txn_id = txn_id.into_inner();
let Some((_, tx)) = self.transactions.remove(&txn_id) else {
return Err(surrealdb_core::rpc::invalid_params("Transaction not found"));
};
tx.commit().await.map_err(surrealdb_core::rpc::types_error_from_anyhow)?;
Ok(DbResult::Other(Value::None))
}
async fn cancel(
&self,
_txn: Option<Uuid>,
_session_id: Option<Uuid>,
params: Array,
) -> Result<DbResult, surrealdb_types::Error> {
let mut params_vec = params.into_vec();
let Some(Value::Uuid(txn_id)) = params_vec.pop() else {
return Err(surrealdb_core::rpc::invalid_params("Expected transaction UUID"));
};
let txn_id = txn_id.into_inner();
let Some((_, tx)) = self.transactions.remove(&txn_id) else {
return Err(surrealdb_core::rpc::invalid_params("Transaction not found"));
};
tx.cancel().await.map_err(surrealdb_core::rpc::types_error_from_anyhow)?;
Ok(DbResult::Other(Value::None))
}
}