pub mod format;
pub mod grpc;
pub mod http;
pub mod response;
pub mod streaming;
pub mod websocket;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use futures::stream::FuturesUnordered;
use surrealdb_core::channel::Receiver;
#[cfg(feature = "graphql")]
use surrealdb_core::graphql::NotificationRouter;
use surrealdb_core::rpc::RpcProtocol;
use surrealdb_rpc::{DbResponse, DbResult};
use surrealdb_types::{Action, Notification};
use tokio::sync::RwLock;
use tokio_stream::StreamExt;
use tokio_util::sync::CancellationToken;
use uuid::Uuid;
#[cfg(feature = "graphql")]
use crate::cnf::GRAPHQL_SUBSCRIPTION_CHANNEL_CAPACITY;
use crate::rpc::websocket::Websocket;
static CONN_CLOSED_ERR: &str = "Connection closed normally";
type WebSocket = Arc<Websocket>;
type WebSockets = RwLock<HashMap<Uuid, WebSocket>>;
#[derive(Clone, Debug)]
pub struct LiveQueryEntry {
pub websocket_id: Uuid,
pub session_id: Uuid,
pub namespace: Option<String>,
pub database: Option<String>,
}
type LiveQueries = RwLock<HashMap<Uuid, LiveQueryEntry>>;
pub struct RpcState {
pub web_sockets: WebSockets,
pub live_queries: LiveQueries,
pub http: Arc<crate::rpc::http::Http>,
pub grpc: Arc<crate::rpc::grpc::Grpc>,
pub metrics_observer: Option<Arc<crate::observe::metrics::MetricsObserver>>,
#[cfg(feature = "graphql")]
pub(crate) notification_router: Arc<NotificationRouter>,
}
impl RpcState {
pub fn new(datastore: Arc<surrealdb_core::kvs::Datastore>) -> Self {
Self::new_with_metrics(datastore, None)
}
pub fn new_with_metrics(
datastore: Arc<surrealdb_core::kvs::Datastore>,
metrics_observer: Option<Arc<crate::observe::metrics::MetricsObserver>>,
) -> Self {
Self::new_with_options(datastore, metrics_observer, None)
}
pub fn new_with_options(
datastore: Arc<surrealdb_core::kvs::Datastore>,
metrics_observer: Option<Arc<crate::observe::metrics::MetricsObserver>>,
durable_session_ttl: Option<std::time::Duration>,
) -> Self {
Self {
web_sockets: RwLock::new(HashMap::new()),
live_queries: RwLock::new(HashMap::new()),
http: Arc::new(crate::rpc::http::Http::new_with_durability(
Arc::clone(&datastore),
durable_session_ttl,
)),
grpc: Arc::new(crate::rpc::grpc::Grpc::new(datastore, metrics_observer.clone())),
metrics_observer,
#[cfg(feature = "graphql")]
notification_router: Arc::new(NotificationRouter::new(
*GRAPHQL_SUBSCRIPTION_CHANNEL_CAPACITY,
)),
}
}
}
pub async fn dispatch_live_notification(notification: Notification, state: Arc<RpcState>) {
#[cfg(feature = "graphql")]
if state.notification_router.has_subscribers() {
state.notification_router.dispatch(¬ification);
}
if state.grpc.dispatch_notification(¬ification).await {
return;
}
let live_query = if notification.action == Action::Killed {
let entry = state.live_queries.write().await.remove(¬ification.id);
if let Some(entry) = entry.as_ref()
&& let Some(obs) = state.metrics_observer.as_ref()
{
obs.adjust_live_query_active(-1, entry.namespace.as_deref(), entry.database.as_deref());
}
entry
} else {
state.live_queries.read().await.get(¬ification.id).cloned()
};
if let Some(entry) = live_query
&& let Some(rpc) = state.web_sockets.read().await.get(&entry.websocket_id).cloned()
{
let wire_session_id = (entry.session_id != rpc.id).then_some(entry.session_id);
let message = DbResponse::success(None, wire_session_id, DbResult::Live(notification));
if rpc.deliver_notification(message) {
if let Some(obs) = state.metrics_observer.as_ref() {
obs.record_live_query_notification(
entry.namespace.as_deref(),
entry.database.as_deref(),
);
}
}
}
}
pub async fn notifications(
channel: Receiver<Notification>,
state: Arc<RpcState>,
canceller: CancellationToken,
) {
let mut futures = FuturesUnordered::new();
loop {
tokio::select! {
biased;
_ = canceller.cancelled() => break,
Some(_) = futures.next() => continue,
Ok(notification) = channel.recv() => {
futures.push(dispatch_live_notification(notification, Arc::clone(&state)));
},
}
}
}
pub async fn graceful_shutdown(state: Arc<RpcState>) {
state.grpc.cleanup_all_lqs().await;
state.grpc.cleanup_all_txns().await;
for (_, rpc) in state.web_sockets.read().await.iter() {
rpc.shutdown.cancel();
}
while !state.web_sockets.read().await.is_empty() {
tokio::time::sleep(Duration::from_millis(250)).await;
}
}
pub fn shutdown(state: &Arc<RpcState>) {
if let Ok(mut writer) = state.web_sockets.try_write() {
writer.drain();
}
}