use std::collections::{HashMap as StdHashMap, HashSet as StdHashSet, VecDeque};
use std::pin::Pin;
use std::sync::Arc;
use std::time::Duration;
use dashmap::DashMap;
use futures::{Stream, StreamExt};
use surrealdb_core::channel::{Receiver, bounded};
use surrealdb_core::ctx::CancelHandle;
use surrealdb_core::dbs::{AuthPrincipalSnapshot, QueryStreamJob, Session};
use surrealdb_core::iam::check::check_ns_db;
use surrealdb_core::kvs::Datastore;
use surrealdb_core::rpc::{RpcProtocol, live_query_owner, types_error_from_anyhow};
use surrealdb_datastore::Transaction;
use surrealdb_iam::Auth;
use surrealdb_kvs::TransactionType;
use surrealdb_protocol::method_names;
use surrealdb_protocol::proto::rpc::v1 as rpc;
use surrealdb_protocol::proto::rpc::v1::surreal_db_service_server::SurrealDbService;
use surrealdb_protocol::proto::v1 as proto;
use surrealdb_rpc::capabilities::{MethodTarget, RouteTarget};
use surrealdb_rpc::error::{invalid_params, method_not_allowed, session_exists, session_not_found};
use surrealdb_rpc::{
DbResult, Method, QUERY_STREAM_BUFFER, QueryResult, QueryStreamItem, QueryType, Token, export,
};
use surrealdb_types::{
Array, Error as TypesError, ErrorDetails as TypesErrorDetails, HashMap, Notification,
QueryError, SurrealValue, Value,
};
use tokio::sync::{RwLock, mpsc};
use tonic::{Request, Response, Status};
use uuid::Uuid;
use web_time::Instant;
use crate::cnf::{
GRPC_MAX_ATTACHED_SESSIONS, GRPC_MAX_MESSAGE_SIZE, GRPC_NOTIFICATION_BUFFER,
HTTP_MAX_IMPORT_BODY_SIZE, MAX_TRANSACTIONS_PER_SESSION, PKG_NAME, PKG_VERSION,
};
use crate::rpc::RpcState;
const EXPORT_CHUNK_SIZE: usize = surrealdb_protocol::DEFAULT_FILE_CHUNK_SIZE;
const _: () = assert!(
EXPORT_CHUNK_SIZE + (crate::cnf::GRPC_MIN_MESSAGE_SIZE / 256) + (1 << 10)
<= crate::cnf::GRPC_MIN_MESSAGE_SIZE,
"an export chunk plus its framing and compression headroom must fit the smallest \
configurable message"
);
const QUERY_BATCH_RECORDS: usize = 256;
fn transport_reserve(limit: usize) -> usize {
(limit / 256) + (1 << 10)
}
const WIRE_FIELD_OVERHEAD: usize = 6;
fn frame_payload_budget() -> usize {
let limit = *GRPC_MAX_MESSAGE_SIZE;
limit.saturating_sub(transport_reserve(limit)).max(1)
}
fn checked_response<M: surrealdb_types::prost::Message>(message: M) -> Result<M, Status> {
let size = message.encoded_len();
let limit = frame_payload_budget();
if size > limit {
return Err(Status::out_of_range(format!(
"The response encodes to {size} bytes, above this server's gRPC message limit of \
{limit} bytes. Raise SURREAL_GRPC_MAX_MESSAGE_SIZE on this server."
)));
}
Ok(message)
}
fn bounded_notification(frame: rpc::subscribe_response::Frame) -> rpc::subscribe_response::Frame {
use surrealdb_types::prost::Message;
let limit = frame_payload_budget();
let response = rpc::SubscribeResponse {
frame: Some(frame),
};
let size = response.encoded_len();
if size <= limit {
return response.frame.expect("the frame this was built around");
}
rpc::subscribe_response::Frame::Error(to_proto_error(&oversized_record(size, limit)))
}
fn record_cost(value: &proto::Value) -> usize {
use surrealdb_types::prost::Message;
value.encoded_len() + WIRE_FIELD_OVERHEAD
}
fn oversized_record(size: usize, budget: usize) -> TypesError {
TypesError::validation(
format!(
"A single record encodes to {size} bytes, above the {budget} bytes one gRPC message \
can carry. Raise SURREAL_GRPC_MAX_MESSAGE_SIZE on this server, or select fewer \
fields."
),
surrealdb_types::ValidationError::InvalidRequest,
)
}
struct LiveQuery {
session_id: Uuid,
namespace: Option<String>,
database: Option<String>,
subscriber: Option<Subscription>,
}
struct Subscription {
id: Uuid,
frames: mpsc::Sender<Result<rpc::SubscribeResponse, Status>>,
}
pub struct Grpc {
kvs: Arc<Datastore>,
sessions: HashMap<Uuid, Arc<RwLock<Session>>>,
ephemeral_sessions: HashMap<Uuid, ()>,
transactions: DashMap<Uuid, (Uuid, Arc<Transaction>)>,
transaction_counts: DashMap<Uuid, usize>,
live_queries: DashMap<Uuid, LiveQuery>,
live_query_counts: DashMap<Uuid, usize>,
metrics_observer: Option<Arc<crate::observe::metrics::MetricsObserver>>,
}
impl Grpc {
pub fn new(
kvs: Arc<Datastore>,
metrics_observer: Option<Arc<crate::observe::metrics::MetricsObserver>>,
) -> Self {
Self {
kvs,
sessions: HashMap::new(),
ephemeral_sessions: HashMap::new(),
transactions: DashMap::new(),
transaction_counts: DashMap::new(),
live_queries: DashMap::new(),
live_query_counts: DashMap::new(),
metrics_observer,
}
}
fn register_ephemeral_session(&self, id: Uuid, session: Arc<RwLock<Session>>) {
self.ephemeral_sessions.insert(id, ());
self.sessions.insert(id, session);
}
fn attached_session_count(&self) -> usize {
self.sessions.len().saturating_sub(self.ephemeral_sessions.len())
}
async fn remove_ephemeral_session(&self, id: &Uuid) {
self.sessions.remove(id);
self.ephemeral_sessions.remove(id);
self.cleanup_lqs(id).await;
self.cleanup_txns(id).await;
}
async fn verify_caller_for_session(
&self,
session_id: &Uuid,
caller: &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 = session_lock.read().await;
if matches!(caller.level(), surrealdb_iam::Level::No)
|| matches!(session.au.level(), surrealdb_iam::Level::No)
{
return Ok(());
}
if session.au.id() == caller.id() && session.au.level() == caller.level() {
Ok(())
} else {
Err(session_not_found(*session_id))
}
}
fn reserve_transaction_slot(&self, session_id: Uuid) -> bool {
let mut count = self.transaction_counts.entry(session_id).or_insert(0);
if *count >= *MAX_TRANSACTIONS_PER_SESSION {
return false;
}
*count += 1;
true
}
fn transaction_belongs_to(&self, txn: &Uuid, session_id: Uuid) -> bool {
self.transactions.get(txn).is_some_and(|entry| entry.value().0 == session_id)
}
fn count_live_query(&self, session_id: Uuid) {
*self.live_query_counts.entry(session_id).or_insert(0) += 1;
}
fn uncount_live_query(&self, session_id: &Uuid) {
if let Some(mut count) = self.live_query_counts.get_mut(session_id)
&& *count > 0
{
*count -= 1;
}
self.live_query_counts.remove_if(session_id, |_, count| *count == 0);
}
fn release_transaction_slot(&self, session_id: &Uuid) {
if let Some(mut count) = self.transaction_counts.get_mut(session_id)
&& *count > 0
{
*count -= 1;
}
self.transaction_counts.remove_if(session_id, |_, count| *count == 0);
}
pub(crate) async fn dispatch_notification(&self, notification: &Notification) -> bool {
let id = notification.id.into_inner();
let Some((subscriber, namespace, database)) = self.live_queries.get(&id).map(|lq| {
(
lq.subscriber.as_ref().map(|s| s.frames.clone()),
lq.namespace.clone(),
lq.database.clone(),
)
}) else {
return false;
};
let Some(subscriber) = subscriber else {
if notification.action == surrealdb_types::Action::Killed {
self.forget_live_query(&id);
}
return true;
};
let frame = match notification.action {
surrealdb_types::Action::Killed => {
rpc::subscribe_response::Frame::End(rpc::SubscribeEnd {
reason: rpc::SubscribeEndReason::Killed as i32,
cursor: None,
})
}
surrealdb_types::Action::Error => rpc::subscribe_response::Frame::Error(bounded_error(
proto::SurrealError::new(
proto::ErrorKind::Query,
notification.result.clone().into_string().unwrap_or_else(|_| {
"The live query raised an evaluation error".to_string()
}),
),
frame_payload_budget(),
)),
action => match to_proto_notification(notification, action) {
Ok(notification) => {
if let Some(observer) = self.metrics_observer.as_ref() {
observer.record_live_query_notification(
namespace.as_deref(),
database.as_deref(),
);
}
rpc::subscribe_response::Frame::Notification(notification)
}
Err(error) => rpc::subscribe_response::Frame::Error(to_proto_error(&error)),
},
};
let frame = bounded_notification(frame);
let mut terminal = !matches!(frame, rpc::subscribe_response::Frame::Notification(_));
if !terminal && subscriber.capacity() <= 1 {
warn!("Ending gRPC subscription to live query {id}: the subscriber is not keeping up");
subscriber
.try_send(Err(Status::resource_exhausted(
"Notifications were produced faster than this subscription read them",
)))
.ok();
terminal = true;
} else {
subscriber
.try_send(Ok(rpc::SubscribeResponse {
frame: Some(frame),
}))
.ok();
}
if terminal {
drop(subscriber);
if notification.action == surrealdb_types::Action::Killed {
self.forget_live_query(&id);
} else if let Some(mut entry) = self.live_queries.get_mut(&id) {
entry.subscriber = None;
}
}
true
}
fn end_subscription(&self, live_query_id: &Uuid, reason: rpc::SubscribeEndReason) {
let subscriber = self
.live_queries
.get(live_query_id)
.and_then(|lq| lq.subscriber.as_ref().map(|s| s.frames.clone()));
if let Some(subscriber) = subscriber {
let frame = rpc::SubscribeResponse {
frame: Some(rpc::subscribe_response::Frame::End(rpc::SubscribeEnd {
reason: reason as i32,
cursor: None,
})),
};
subscriber.try_send(Ok(frame)).ok();
}
}
async fn discard_live_query(&self, live_query_id: &Uuid) {
self.end_subscription(live_query_id, rpc::SubscribeEndReason::Killed);
self.forget_live_query(live_query_id);
if let Err(err) = self.kvs.delete_queries(vec![*live_query_id]).await {
error!("Error discarding live query {live_query_id}: {err}");
}
}
fn forget_live_query(&self, live_query_id: &Uuid) -> Option<LiveQuery> {
let (_, entry) = self.live_queries.remove(live_query_id)?;
self.uncount_live_query(&entry.session_id);
if let Some(observer) = self.metrics_observer.as_ref() {
observer.adjust_live_query_active(
-1,
entry.namespace.as_deref(),
entry.database.as_deref(),
);
}
Some(entry)
}
pub(crate) async fn cleanup_all_txns(&self) {
self.cleanup_txns_filtered(None).await;
}
async fn cleanup_txns_filtered(&self, session_filter: Option<&Uuid>) {
if let Some(session_id) = session_filter
&& !self.transaction_counts.contains_key(session_id)
{
return;
}
let doomed: Vec<Uuid> = self
.transactions
.iter()
.filter(|entry| match session_filter {
Some(session_id) => &entry.value().0 == session_id,
None => true,
})
.map(|entry| *entry.key())
.collect();
for id in doomed {
if let Some((_, (session_id, tx))) = self.transactions.remove(&id) {
self.release_transaction_slot(&session_id);
if let Err(err) = tx.cancel().await {
warn!("Error cancelling gRPC transaction {id}: {err}");
}
}
}
}
}
impl RpcProtocol for Grpc {
fn kvs(&self) -> &Datastore {
&self.kvs
}
fn kvs_arc(&self) -> Arc<Datastore> {
Arc::clone(&self.kvs)
}
fn version_data(&self) -> DbResult {
DbResult::Other(Value::String(format!("{PKG_NAME}-{}", *PKG_VERSION)))
}
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.sessions.contains_key(&session_id) {
return Err(session_exists(session_id));
}
if self.attached_session_count() >= *GRPC_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);
self.sessions.insert(session_id, Arc::new(RwLock::new(session)));
Ok(DbResult::Other(Value::None))
}
async fn get_tx(&self, id: Uuid) -> Result<Arc<Transaction>, TypesError> {
self.transactions
.get(&id)
.map(|entry| Arc::clone(&entry.value().1))
.ok_or_else(|| invalid_params("Transaction not found"))
}
const LQ_SUPPORT: bool = true;
async fn handle_live(
&self,
lqid: &Uuid,
session_id: Uuid,
namespace: Option<String>,
database: Option<String>,
) {
self.count_live_query(session_id);
self.live_queries.insert(
*lqid,
LiveQuery {
session_id,
namespace: namespace.clone(),
database: database.clone(),
subscriber: None,
},
);
if let Some(observer) = self.metrics_observer.as_ref() {
observer.adjust_live_query_active(1, namespace.as_deref(), database.as_deref());
}
if !self.sessions.contains_key(&session_id) {
self.discard_live_query(lqid).await;
return;
}
trace!("Registered live query {lqid} on the gRPC transport");
}
async fn cleanup_lqs(&self, session_id: &Uuid) {
if !self.live_query_counts.contains_key(session_id) {
return;
}
let doomed: Vec<Uuid> = self
.live_queries
.iter()
.filter(|entry| &entry.value().session_id == session_id)
.map(|entry| *entry.key())
.collect();
for id in doomed {
self.end_subscription(&id, rpc::SubscribeEndReason::SessionClosed);
self.forget_live_query(&id);
if let Err(err) = self.kvs.delete_queries(vec![id]).await {
error!("Error cleaning up live query {id} on the gRPC transport: {err}");
}
}
}
async fn cleanup_all_lqs(&self) {
let doomed: Vec<Uuid> = self.live_queries.iter().map(|entry| *entry.key()).collect();
for id in doomed {
self.end_subscription(&id, rpc::SubscribeEndReason::ServerShutdown);
self.forget_live_query(&id);
if let Err(err) = self.kvs.delete_queries(vec![id]).await {
error!("Error cleaning up live query {id} on shutdown: {err}");
}
}
}
async fn cleanup_txns(&self, session_id: &Uuid) {
self.cleanup_txns_filtered(Some(session_id)).await;
}
async fn begin(&self, _txn: Option<Uuid>, session_id: Uuid) -> Result<DbResult, TypesError> {
self.get_session(&session_id).await?;
if self.ephemeral_sessions.contains_key(&session_id) {
return Err(invalid_params("Opening a transaction requires an attached session"));
}
if !self.reserve_transaction_slot(session_id) {
return Err(surrealdb_rpc::error::too_many_transactions());
}
let tx = match self.kvs.transaction(TransactionType::Write).await {
Ok(tx) => tx,
Err(err) => {
self.release_transaction_slot(&session_id);
return Err(types_error_from_anyhow(err));
}
};
let id = Uuid::now_v7();
self.transactions.insert(id, (session_id, Arc::new(tx)));
if !self.sessions.contains_key(&session_id) {
self.cleanup_txns_filtered(Some(&session_id)).await;
return Err(session_not_found(session_id));
}
Ok(DbResult::Other(Value::Uuid(surrealdb_types::Uuid::from(id))))
}
async fn commit(
&self,
txn: Option<Uuid>,
_session_id: Uuid,
_params: Array,
) -> Result<DbResult, TypesError> {
let txn = txn.ok_or_else(|| invalid_params("Expected a transaction id"))?;
let Some((_, (session_id, tx))) = self.transactions.remove(&txn) else {
return Err(invalid_params("Transaction not found"));
};
self.release_transaction_slot(&session_id);
tx.commit().await.map_err(types_error_from_anyhow)?;
Ok(DbResult::Other(Value::None))
}
async fn cancel(
&self,
txn: Option<Uuid>,
_session_id: Uuid,
_params: Array,
) -> Result<DbResult, TypesError> {
let txn = txn.ok_or_else(|| invalid_params("Expected a transaction id"))?;
let Some((_, (session_id, tx))) = self.transactions.remove(&txn) else {
return Err(invalid_params("Transaction not found"));
};
self.release_transaction_slot(&session_id);
tx.cancel().await.map_err(types_error_from_anyhow)?;
Ok(DbResult::Other(Value::None))
}
}
#[derive(Clone, Copy)]
struct ResolvedSession {
id: Uuid,
client: Option<Uuid>,
}
pub struct GrpcService {
state: Arc<RpcState>,
caller: Session,
}
impl GrpcService {
pub fn new(state: Arc<RpcState>, caller: Session) -> Self {
Self {
state,
caller,
}
}
fn rpc(&self) -> &Grpc {
&self.state.grpc
}
fn kvs(&self) -> &Datastore {
&self.rpc().kvs
}
async fn resolve(
&self,
context: Option<&rpc::RequestContext>,
gate: bool,
) -> Result<ResolvedSession, Status> {
match context.and_then(|context| context.session.as_ref()) {
Some(session) => {
let id = to_uuid(session)?;
if gate {
self.rpc()
.verify_caller_for_session(&id, self.caller.au.as_ref())
.await
.map_err(|err| to_status(&err))?;
}
Ok(ResolvedSession {
id,
client: Some(id),
})
}
None => {
let id = Uuid::new_v4();
let mut session = self.caller.clone();
session.id = Some(id);
self.rpc().register_ephemeral_session(id, Arc::new(RwLock::new(session)));
Ok(ResolvedSession {
id,
client: None,
})
}
}
}
async fn release(&self, session: &ResolvedSession) {
if session.client.is_none() {
self.rpc().remove_ephemeral_session(&session.id).await;
}
}
async fn execute(
&self,
context: Option<&rpc::RequestContext>,
method: Method,
params: Array,
) -> Result<DbResult, Status> {
let txn = match context.and_then(|context| context.transaction.as_ref()) {
Some(txn) => Some(to_uuid(txn)?),
None => None,
};
let session = self.resolve(context, method != Method::Attach).await?;
if let Some(txn) = txn
&& !self.rpc().transaction_belongs_to(&txn, session.id)
{
self.release(&session).await;
return Err(to_status(&invalid_params("Transaction not found")));
}
let dispatch =
RpcProtocol::execute(self.rpc(), txn, session.id, session.client, method, params);
let deadline = context
.and_then(|context| context.timeout.as_ref())
.and_then(|timeout| Duration::try_from(*timeout).ok())
.filter(|_| method == Method::Query);
let result = match deadline {
Some(deadline) => match tokio::time::timeout(deadline, dispatch).await {
Ok(result) => result,
Err(_) => Err(surrealdb_core::rpc::query_timeout_error(deadline)),
},
None => dispatch.await,
};
self.release(&session).await;
result.map_err(|err| to_status(&err))
}
async fn session_for(
&self,
context: Option<&rpc::RequestContext>,
route: RouteTarget,
) -> Result<Session, Status> {
if !self.kvs().allows_http_route(&route) {
warn!("Capabilities denied gRPC route request attempt, target: '{route}'");
return Err(Status::permission_denied(format!("Route {route} is not allowed")));
}
let resolved = self.resolve(context, true).await?;
let session = match self.rpc().get_session(&resolved.id).await {
Ok(lock) => lock.read().await.clone(),
Err(err) => {
self.release(&resolved).await;
return Err(to_status(&err));
}
};
self.release(&resolved).await;
Ok(session)
}
}
type ResponseStream<T> = Pin<Box<dyn Stream<Item = Result<T, Status>> + Send>>;
#[tonic::async_trait]
impl SurrealDbService for GrpcService {
async fn get_capabilities(
&self,
request: Request<rpc::GetCapabilitiesRequest>,
) -> Result<Response<rpc::GetCapabilitiesResponse>, Status> {
let request = request.into_inner();
if let Some(client) = request.client.as_ref() {
debug!(
"gRPC client connected: {} {} on {}",
client.name, client.version, client.platform
);
}
Ok(Response::new(checked_response(rpc::GetCapabilitiesResponse {
capabilities: Some(self.server_capabilities()),
})?))
}
async fn health(
&self,
_request: Request<rpc::HealthRequest>,
) -> Result<Response<rpc::HealthResponse>, Status> {
if !self.kvs().allows_http_route(&RouteTarget::Health) {
return Err(Status::permission_denied("Route health is not allowed"));
}
self.kvs().health_check().await.map_err(|err| {
error!("Health check failed: {err}");
Status::unavailable("Health check failed")
})?;
Ok(Response::new(checked_response(rpc::HealthResponse {})?))
}
async fn attach_session(
&self,
request: Request<rpc::AttachSessionRequest>,
) -> Result<Response<rpc::AttachSessionResponse>, Status> {
let context = request.into_inner().context;
if let Some(session) = context.as_ref().and_then(|context| context.session.as_ref()) {
let id = to_uuid(session)?;
self.rpc()
.verify_caller_for_session(&id, self.caller.au.as_ref())
.await
.map_err(|err| to_status(&err))?;
return Ok(Response::new(checked_response(rpc::AttachSessionResponse {
session: Some(proto::Uuid::from_uuid(id)),
created: false,
})?));
}
let id = Uuid::new_v4();
let context = rpc::RequestContext {
session: Some(proto::Uuid::from_uuid(id)),
transaction: None,
timeout: context.and_then(|context| context.timeout),
};
self.execute(Some(&context), Method::Attach, Array::new()).await?;
Ok(Response::new(checked_response(rpc::AttachSessionResponse {
session: Some(proto::Uuid::from_uuid(id)),
created: true,
})?))
}
async fn detach_session(
&self,
request: Request<rpc::DetachSessionRequest>,
) -> Result<Response<rpc::DetachSessionResponse>, Status> {
let context = request.into_inner().context;
self.execute(context.as_ref(), Method::Detach, Array::new()).await?;
Ok(Response::new(checked_response(rpc::DetachSessionResponse {})?))
}
async fn reset_session(
&self,
request: Request<rpc::ResetSessionRequest>,
) -> Result<Response<rpc::ResetSessionResponse>, Status> {
let context = request.into_inner().context;
self.execute(context.as_ref(), Method::Reset, Array::new()).await?;
Ok(Response::new(checked_response(rpc::ResetSessionResponse {})?))
}
async fn r#use(
&self,
request: Request<rpc::UseRequest>,
) -> Result<Response<rpc::UseResponse>, Status> {
let request = request.into_inner();
let params =
Array::from(vec![from_nullable(request.namespace), from_nullable(request.database)]);
let result = self.execute(request.context.as_ref(), Method::Use, params).await?;
let selection = match result {
DbResult::Other(Value::Object(object)) => object,
_ => {
return Err(Status::internal("Use did not report the resulting selection"));
}
};
let field = |key: &str| match selection.get(key) {
Some(Value::String(value)) => value.clone(),
_ => String::new(),
};
Ok(Response::new(checked_response(rpc::UseResponse {
namespace: field("namespace"),
database: field("database"),
})?))
}
async fn set_variable(
&self,
request: Request<rpc::SetVariableRequest>,
) -> Result<Response<rpc::SetVariableResponse>, Status> {
let request = request.into_inner();
let value = match request.value {
Some(value) => from_proto_value(value)?,
None => Value::None,
};
let params = Array::from(vec![Value::String(request.name), value]);
self.execute(request.context.as_ref(), Method::Set, params).await?;
Ok(Response::new(checked_response(rpc::SetVariableResponse {})?))
}
async fn unset_variable(
&self,
request: Request<rpc::UnsetVariableRequest>,
) -> Result<Response<rpc::UnsetVariableResponse>, Status> {
let request = request.into_inner();
let params = Array::from(vec![Value::String(request.name)]);
self.execute(request.context.as_ref(), Method::Unset, params).await?;
Ok(Response::new(checked_response(rpc::UnsetVariableResponse {})?))
}
async fn signup(
&self,
request: Request<rpc::SignupRequest>,
) -> Result<Response<rpc::SignupResponse>, Status> {
let request = request.into_inner();
let credentials = request
.credentials
.ok_or_else(|| Status::invalid_argument("Expected signup credentials"))?;
let params = Array::from(vec![Value::Object(record_credentials(credentials)?)]);
let result = self.execute(request.context.as_ref(), Method::Signup, params).await?;
Ok(Response::new(checked_response(rpc::SignupResponse {
tokens: Some(to_tokens(result)?),
})?))
}
async fn signin(
&self,
request: Request<rpc::SigninRequest>,
) -> Result<Response<rpc::SigninResponse>, Status> {
let request = request.into_inner();
let access = request
.access_method
.ok_or_else(|| Status::invalid_argument("Expected an access method"))?;
let params = Array::from(vec![Value::Object(access_credentials(access)?)]);
let result = self.execute(request.context.as_ref(), Method::Signin, params).await?;
Ok(Response::new(checked_response(rpc::SigninResponse {
tokens: Some(to_tokens(result)?),
})?))
}
async fn authenticate(
&self,
request: Request<rpc::AuthenticateRequest>,
) -> Result<Response<rpc::AuthenticateResponse>, Status> {
let request = request.into_inner();
let params = Array::from(vec![Value::String(request.token)]);
let result = self.execute(request.context.as_ref(), Method::Authenticate, params).await?;
Ok(Response::new(checked_response(rpc::AuthenticateResponse {
expires_at: None,
tokens: to_tokens(result).ok(),
})?))
}
async fn refresh_tokens(
&self,
request: Request<rpc::RefreshTokensRequest>,
) -> Result<Response<rpc::RefreshTokensResponse>, Status> {
let request = request.into_inner();
let params = Array::from(vec![token_value(request.access, request.refresh)]);
let result = self.execute(request.context.as_ref(), Method::Refresh, params).await?;
Ok(Response::new(checked_response(rpc::RefreshTokensResponse {
tokens: Some(to_tokens(result)?),
})?))
}
async fn revoke_tokens(
&self,
request: Request<rpc::RevokeTokensRequest>,
) -> Result<Response<rpc::RevokeTokensResponse>, Status> {
let request = request.into_inner();
let params = Array::from(vec![token_value(request.access, request.refresh)]);
self.execute(request.context.as_ref(), Method::Revoke, params).await?;
Ok(Response::new(checked_response(rpc::RevokeTokensResponse {})?))
}
async fn invalidate(
&self,
request: Request<rpc::InvalidateRequest>,
) -> Result<Response<rpc::InvalidateResponse>, Status> {
let context = request.into_inner().context;
self.execute(context.as_ref(), Method::Invalidate, Array::new()).await?;
Ok(Response::new(checked_response(rpc::InvalidateResponse {})?))
}
async fn begin_transaction(
&self,
request: Request<rpc::BeginTransactionRequest>,
) -> Result<Response<rpc::BeginTransactionResponse>, Status> {
let context = request.into_inner().context;
let result = self.execute(context.as_ref(), Method::Begin, Array::new()).await?;
let DbResult::Other(Value::Uuid(id)) = result else {
return Err(Status::internal("Begin did not return a transaction id"));
};
Ok(Response::new(checked_response(rpc::BeginTransactionResponse {
transaction: Some(proto::Uuid::from_uuid(id.into_inner())),
})?))
}
async fn commit_transaction(
&self,
request: Request<rpc::CommitTransactionRequest>,
) -> Result<Response<rpc::CommitTransactionResponse>, Status> {
let context = request.into_inner().context;
self.execute(context.as_ref(), Method::Commit, Array::new()).await?;
Ok(Response::new(checked_response(rpc::CommitTransactionResponse {})?))
}
async fn cancel_transaction(
&self,
request: Request<rpc::CancelTransactionRequest>,
) -> Result<Response<rpc::CancelTransactionResponse>, Status> {
let context = request.into_inner().context;
self.execute(context.as_ref(), Method::Cancel, Array::new()).await?;
Ok(Response::new(checked_response(rpc::CancelTransactionResponse {})?))
}
type QueryStream = ResponseStream<rpc::QueryResponse>;
async fn query(
&self,
request: Request<rpc::QueryRequest>,
) -> Result<Response<Self::QueryStream>, Status> {
let request = request.into_inner();
let arrow_only = !request.accepted_encodings.is_empty()
&& request
.accepted_encodings
.iter()
.all(|encoding| *encoding == rpc::ResultEncoding::Arrow as i32);
if arrow_only {
return Err(Status::unimplemented(
"This server does not serve columnar (Arrow) query results",
));
}
let batch_records = batch_records(request.max_batch_records);
self.stream_query(request.context.as_ref(), request.query, request.variables, batch_records)
.await
}
async fn run(
&self,
request: Request<rpc::RunRequest>,
) -> Result<Response<rpc::RunResponse>, Status> {
let request = request.into_inner();
let args = request
.args
.into_iter()
.map(from_proto_value)
.collect::<Result<Vec<Value>, Status>>()?;
let version = if request.version.is_empty() {
Value::None
} else {
Value::String(request.version)
};
let params = Array::from(vec![
Value::String(request.name),
version,
Value::Array(Array::from(args)),
]);
let result = self.execute(request.context.as_ref(), Method::Run, params).await?;
let DbResult::Other(value) = result else {
return Err(Status::internal("Run did not return a value"));
};
Ok(Response::new(checked_response(rpc::RunResponse {
result: Some(to_proto_value(value)?),
})?))
}
async fn kill(
&self,
request: Request<rpc::KillRequest>,
) -> Result<Response<rpc::KillResponse>, Status> {
let request = request.into_inner();
let live_query_id = request
.live_query_id
.ok_or_else(|| Status::invalid_argument("Expected a live query id"))?;
let params = Array::from(vec![Value::Uuid(to_uuid(&live_query_id)?.into())]);
self.execute(request.context.as_ref(), Method::Kill, params).await?;
Ok(Response::new(checked_response(rpc::KillResponse {})?))
}
type SubscribeStream = ResponseStream<rpc::SubscribeResponse>;
async fn subscribe(
&self,
request: Request<rpc::SubscribeRequest>,
) -> Result<Response<Self::SubscribeStream>, Status> {
let request = request.into_inner();
if request.resume_from.is_some() {
return Err(Status::unimplemented(
"This server does not retain live query history to resume from",
));
}
let session = self.resolve(request.context.as_ref(), true).await?;
let Some(session_id) = session.client else {
self.release(&session).await;
return Err(rejected_before_execution("Subscribing requires an attached session"));
};
let (live_query_id, owned) = match request.subscribe_to {
Some(rpc::subscribe_request::SubscribeTo::LiveQueryId(id)) => (to_uuid(&id)?, false),
Some(rpc::subscribe_request::SubscribeTo::Query(registration)) => {
let id = self.register_live_query(request.context.as_ref(), registration).await?;
(id, true)
}
None => {
return Err(Status::invalid_argument("Expected a live query id or a query"));
}
};
self.attach_subscription(session_id, live_query_id, owned)
}
async fn import_surql(
&self,
request: Request<tonic::Streaming<rpc::ImportSurqlRequest>>,
) -> Result<Response<rpc::ImportSurqlResponse>, Status> {
self.run_import(request.into_inner()).await
}
async fn import_ml_model(
&self,
request: Request<tonic::Streaming<rpc::ImportMlModelRequest>>,
) -> Result<Response<rpc::ImportMlModelResponse>, Status> {
self.run_ml_import(request.into_inner()).await
}
type ExportSurqlStream = ResponseStream<rpc::ExportSurqlResponse>;
async fn export_surql(
&self,
request: Request<rpc::ExportSurqlRequest>,
) -> Result<Response<Self::ExportSurqlStream>, Status> {
let request = request.into_inner();
let session = self.session_for(request.context.as_ref(), RouteTarget::Export).await?;
let config = match request.config {
Some(config) => from_proto_export_config(config),
None => export::Config::default(),
};
let export = self.start_export(session, config).await?;
Ok(Response::new(Box::pin(frame_byte_stream(export, |frame| rpc::ExportSurqlResponse {
frame: Some(match frame {
ByteFrame::Chunk(chunk) => rpc::export_surql_response::Frame::Chunk(chunk),
ByteFrame::Trailer(trailer) => rpc::export_surql_response::Frame::Trailer(trailer),
ByteFrame::Error(error) => rpc::export_surql_response::Frame::Error(error),
}),
}))))
}
type ExportDirectoryStream = ResponseStream<rpc::ExportDirectoryResponse>;
async fn export_directory(
&self,
_request: Request<rpc::ExportDirectoryRequest>,
) -> Result<Response<Self::ExportDirectoryStream>, Status> {
Err(Status::unimplemented(
"This server does not produce directory-format exports; use ExportSurql",
))
}
type ExportMlModelStream = ResponseStream<rpc::ExportMlModelResponse>;
async fn export_ml_model(
&self,
request: Request<rpc::ExportMlModelRequest>,
) -> Result<Response<Self::ExportMlModelStream>, Status> {
let request = request.into_inner();
let session = self.session_for(request.context.as_ref(), RouteTarget::Ml).await?;
let export = self.start_ml_export(session, request.name, request.version).await?;
Ok(Response::new(Box::pin(frame_byte_stream(export, |frame| rpc::ExportMlModelResponse {
frame: Some(match frame {
ByteFrame::Chunk(chunk) => rpc::export_ml_model_response::Frame::Chunk(chunk),
ByteFrame::Trailer(trailer) => {
rpc::export_ml_model_response::Frame::Trailer(trailer)
}
ByteFrame::Error(error) => rpc::export_ml_model_response::Frame::Error(error),
}),
}))))
}
}
const ROUTE_METHODS: &[(RouteTarget, &str)] = &[
(RouteTarget::Health, method_names::HEALTH),
(RouteTarget::Export, method_names::EXPORT_SURQL),
(RouteTarget::Import, method_names::IMPORT_SURQL),
(RouteTarget::Ml, method_names::EXPORT_ML_MODEL),
(RouteTarget::Ml, method_names::IMPORT_ML_MODEL),
];
const RPC_METHODS: &[(Method, &str)] = &[
(Method::Use, method_names::USE),
(Method::Set, method_names::SET_VARIABLE),
(Method::Unset, method_names::UNSET_VARIABLE),
(Method::Signup, method_names::SIGNUP),
(Method::Signin, method_names::SIGNIN),
(Method::Authenticate, method_names::AUTHENTICATE),
(Method::Refresh, method_names::REFRESH_TOKENS),
(Method::Revoke, method_names::REVOKE_TOKENS),
(Method::Invalidate, method_names::INVALIDATE),
(Method::Begin, method_names::BEGIN_TRANSACTION),
(Method::Commit, method_names::COMMIT_TRANSACTION),
(Method::Cancel, method_names::CANCEL_TRANSACTION),
(Method::Query, method_names::QUERY),
(Method::Run, method_names::RUN),
(Method::Kill, method_names::KILL),
(Method::Attach, method_names::ATTACH_SESSION),
(Method::Detach, method_names::DETACH_SESSION),
(Method::Reset, method_names::RESET_SESSION),
];
impl GrpcService {
fn server_capabilities(&self) -> rpc::ServerCapabilities {
let mut capabilities = vec![
"SESSIONS".to_string(),
"TRANSACTIONS".to_string(),
"LIVE_QUERIES".to_string(),
"REFRESH_TOKENS".to_string(),
];
if cfg!(feature = "ml") {
capabilities.push("ML_MODELS".to_string());
}
let mut denied = Vec::new();
for (route, method) in ROUTE_METHODS {
if !self.kvs().allows_http_route(route) {
denied.push((*method).to_string());
}
}
for (rpc_method, method) in RPC_METHODS {
if !self.kvs().allows_rpc_method(&MethodTarget {
method: *rpc_method,
}) {
denied.push((*method).to_string());
}
}
let version_denied = !self.kvs().allows_rpc_method(&MethodTarget {
method: Method::Version,
});
if version_denied {
denied.push(method_names::VERSION_PSEUDO_METHOD.to_string());
}
rpc::ServerCapabilities {
server_version: if version_denied {
String::new()
} else {
format!("{PKG_NAME}-{}", *PKG_VERSION)
},
low_api_version: None,
high_api_version: None,
capabilities,
denied_methods: denied,
limits: Some(rpc::Limits {
max_message_bytes: *GRPC_MAX_MESSAGE_SIZE as u64,
max_chunk_bytes: EXPORT_CHUNK_SIZE as u64,
max_query_duration: self
.kvs()
.query_timeout()
.and_then(|timeout| proto::Duration::try_from(timeout).ok()),
max_batch_records: QUERY_BATCH_RECORDS as u32,
}),
live_queries: Some(rpc::LiveQueryCapabilities {
delivery: rpc::LiveQueryDelivery::AtMostOnce as i32,
resumable: false,
multiple_subscribers: false,
retention: None,
}),
accepted_message_encodings: crate::ntw::grpc::accepted_message_encodings(),
}
}
async fn run_query(
&self,
context: Option<&rpc::RequestContext>,
query: String,
variables: Option<proto::Variables>,
) -> Result<Vec<QueryResult>, Status> {
let variables = match variables {
Some(variables) => from_proto_variables(variables)?,
None => Value::None,
};
let params = Array::from(vec![Value::String(query), variables]);
match self.execute(context, Method::Query, params).await? {
DbResult::Query(results) => Ok(results),
_ => Err(Status::internal("Query did not return statement results")),
}
}
async fn stream_query(
&self,
context: Option<&rpc::RequestContext>,
query: String,
variables: Option<proto::Variables>,
batch_records: usize,
) -> Result<Response<ResponseStream<rpc::QueryResponse>>, Status> {
let txn = match context.and_then(|context| context.transaction.as_ref()) {
Some(txn) => Some(to_uuid(txn)?),
None => None,
};
let session = self.resolve(context, true).await?;
if let Some(txn) = txn
&& !self.rpc().transaction_belongs_to(&txn, session.id)
{
self.release(&session).await;
return Err(to_status(&invalid_params("Transaction not found")));
}
let variables = match variables {
Some(variables) => match from_proto_variables(variables) {
Ok(variables) => variables,
Err(status) => {
self.release(&session).await;
return Err(status);
}
},
None => Value::None,
};
let params = Array::from(vec![Value::String(query), variables]);
let (items_tx, items_rx) = bounded(QUERY_STREAM_BUFFER);
let cancel = CancelHandle::new();
let requested_timeout = context
.and_then(|context| context.timeout)
.and_then(|timeout| Duration::try_from(timeout).ok());
let (job, principal) = match RpcProtocol::query_stream(
self.rpc(),
txn,
session.id,
params,
Some(cancel.clone()),
items_tx,
)
.await
{
Ok(job) => job,
Err(error) => {
self.release(&session).await;
return Err(to_status(&error));
}
};
let begin = rpc::QueryResponse {
frame: Some(rpc::query_response::Frame::Begin(rpc::QueryBegin {
query_id: Some(proto::Uuid::from_uuid(Uuid::new_v4())),
statement_count: job.statement_count as u32,
})),
};
if let Some(deadline) = shortest(requested_timeout, self.rpc().kvs().query_timeout()) {
let cancel = cancel.clone();
let items = items_rx.clone();
tokio::spawn(async move {
tokio::time::sleep(deadline).await;
cancel.trip();
items.close();
});
}
let state = QueryFraming::new(
Arc::clone(&self.state),
session,
principal,
job,
items_rx,
batch_records,
cancel,
);
let frames = futures::stream::unfold(Some(state), |state| async move {
let mut state = state?;
state.next().await.map(|frame| (Ok(frame), Some(state)))
});
Ok(Response::new(Box::pin(futures::stream::once(async move { Ok(begin) }).chain(frames))))
}
async fn register_live_query(
&self,
context: Option<&rpc::RequestContext>,
registration: rpc::LiveQueryRegistration,
) -> Result<Uuid, Status> {
let results = self.run_query(context, registration.query, registration.variables).await?;
let registered: Vec<Uuid> = results
.iter()
.filter(|result| result.query_type == QueryType::Live)
.filter_map(|result| match &result.result {
Ok(Value::Uuid(id)) => Some(id.into_inner()),
_ => None,
})
.collect();
let outcome = single_live_query(results);
let kept = outcome.as_ref().ok().copied();
for id in registered {
if Some(id) != kept {
self.rpc().discard_live_query(&id).await;
}
}
outcome
}
fn attach_subscription(
&self,
session_id: Uuid,
live_query_id: Uuid,
owned: bool,
) -> Result<Response<ResponseStream<rpc::SubscribeResponse>>, Status> {
let (sender, receiver) = mpsc::channel(*GRPC_NOTIFICATION_BUFFER + 1);
let subscription_id = Uuid::new_v4();
{
let Some(mut entry) = self.rpc().live_queries.get_mut(&live_query_id) else {
return Err(Status::not_found(format!("Live query {live_query_id} not found")));
};
if entry.session_id != session_id {
return Err(Status::not_found(format!("Live query {live_query_id} not found")));
}
if entry.subscriber.is_some() {
return Err(Status::already_exists(format!(
"Live query {live_query_id} already has a subscriber"
)));
}
entry.subscriber = Some(Subscription {
id: subscription_id,
frames: sender,
});
}
let begin = rpc::SubscribeResponse {
frame: Some(rpc::subscribe_response::Frame::Begin(rpc::SubscribeBegin {
subscription_id: Some(proto::Uuid::from_uuid(subscription_id)),
live_query_id: Some(proto::Uuid::from_uuid(live_query_id)),
cursor: None,
})),
};
let guard = SubscriptionGuard {
state: Arc::clone(&self.state),
live_query_id,
subscription_id,
owned,
};
let notifications =
futures::stream::unfold((receiver, guard), |(mut receiver, guard)| async move {
receiver.recv().await.map(|item| (item, (receiver, guard)))
});
let stream = futures::stream::once(async move { Ok(begin) }).chain(notifications);
Ok(Response::new(Box::pin(stream)))
}
async fn run_import(
&self,
mut frames: tonic::Streaming<rpc::ImportSurqlRequest>,
) -> Result<Response<rpc::ImportSurqlResponse>, Status> {
use futures::StreamExt;
let Some(first) = frames.next().await.transpose()? else {
return Err(Status::invalid_argument("Import stream carried no frames"));
};
let Some(rpc::import_surql_request::Frame::Begin(begin)) = first.frame else {
return Err(Status::invalid_argument("Import stream must open with a begin frame"));
};
let session = self.session_for(begin.context.as_ref(), RouteTarget::Import).await?;
self.kvs()
.check(
&session,
surrealdb_iam::Action::Edit,
surrealdb_iam::ResourceKind::Any.on_level(session.au.level().to_owned()),
)
.map_err(|err| Status::permission_denied(err.to_string()))?;
let limit = *HTTP_MAX_IMPORT_BODY_SIZE as u64;
let (bytes_tx, bytes_rx) = surrealdb::channel::bounded::<anyhow::Result<bytes::Bytes>>(1);
let trailer = tokio::spawn(async move {
let mut hasher = blake3::Hasher::new();
let mut streamed: u64 = 0;
while let Some(frame) = frames.next().await {
let frame = match frame {
Ok(frame) => frame,
Err(status) => {
bytes_tx.send(Err(anyhow::anyhow!("{}", status.message()))).await.ok();
return Err(status);
}
};
match frame.frame {
Some(rpc::import_surql_request::Frame::Chunk(chunk)) => {
streamed += chunk.data.len() as u64;
if streamed > limit {
let refused = Status::resource_exhausted(format!(
"Import exceeds the {limit} byte limit"
));
bytes_tx.send(Err(anyhow::anyhow!("{}", refused.message()))).await.ok();
return Err(refused);
}
hasher.update(&chunk.data);
if bytes_tx.send(Ok(chunk.data)).await.is_err() {
return Ok(());
}
}
Some(rpc::import_surql_request::Frame::Trailer(trailer)) => {
drop(bytes_tx);
return verify_trailer(&trailer, streamed, hasher.finalize());
}
Some(rpc::import_surql_request::Frame::Begin(_)) => {
return Err(Status::invalid_argument(
"Import stream carried a second begin frame",
));
}
None => {}
}
}
Err(Status::data_loss(
"Import stream ended without a trailer; the import may be partially applied",
))
});
let result = self.kvs().import_stream(&session, bytes_rx).await;
match trailer.await {
Ok(Ok(())) => {}
Ok(Err(status)) => return Err(status),
Err(err) => return Err(Status::internal(format!("Import framing task failed: {err}"))),
}
let results = result.map_err(|err| Status::invalid_argument(err.to_string()))?;
if let Some(status) = import_failure_status(&results) {
return Err(status);
}
Ok(Response::new(checked_response(rpc::ImportSurqlResponse {})?))
}
#[cfg(feature = "ml")]
async fn run_ml_import(
&self,
mut frames: tonic::Streaming<rpc::ImportMlModelRequest>,
) -> Result<Response<rpc::ImportMlModelResponse>, Status> {
use futures::StreamExt;
let Some(first) = frames.next().await.transpose()? else {
return Err(Status::invalid_argument("Import stream carried no frames"));
};
let Some(rpc::import_ml_model_request::Frame::Begin(begin)) = first.frame else {
return Err(Status::invalid_argument("Import stream must open with a begin frame"));
};
let session = self.session_for(begin.context.as_ref(), RouteTarget::Ml).await?;
let (namespace, database) =
check_ns_db(&session).map_err(|err| rejected_before_execution(err.to_string()))?;
self.kvs()
.check(
&session,
surrealdb_iam::Action::Edit,
surrealdb_iam::ResourceKind::Model.on_db(&namespace, &database),
)
.map_err(|err| Status::permission_denied(err.to_string()))?;
let limit = *HTTP_MAX_IMPORT_BODY_SIZE;
let mut hasher = blake3::Hasher::new();
let mut model = Vec::new();
let mut trailer = None;
while let Some(frame) = frames.next().await {
match frame?.frame {
Some(rpc::import_ml_model_request::Frame::Chunk(chunk)) => {
if model.len() + chunk.data.len() > limit {
return Err(Status::resource_exhausted(format!(
"Import exceeds the {limit} byte limit"
)));
}
hasher.update(&chunk.data);
model.extend_from_slice(&chunk.data);
}
Some(rpc::import_ml_model_request::Frame::Trailer(sent)) => {
trailer = Some(sent);
break;
}
Some(rpc::import_ml_model_request::Frame::Begin(_)) => {
return Err(Status::invalid_argument(
"Import stream carried a second begin frame",
));
}
None => {}
}
}
let trailer = trailer.ok_or_else(|| {
Status::data_loss("Import stream ended without a trailer; the model was not stored")
})?;
verify_trailer(&trailer, model.len() as u64, hasher.finalize())?;
let file = surrealml_core::storage::surml_file::SurMlFile::from_bytes(model)
.map_err(|err| Status::invalid_argument(format!("Invalid SurrealML file: {err}")))?;
let (name, version) = (file.header.name.to_string(), file.header.version.to_string());
if name.is_empty() || version.is_empty() {
return Err(Status::invalid_argument("Model name and version must be set"));
}
if !begin.name.is_empty() && begin.name != name {
return Err(Status::invalid_argument(format!(
"The request names model {} but the file carries {name}",
begin.name
)));
}
if !begin.version.is_empty() && begin.version != version {
return Err(Status::invalid_argument(format!(
"The request names version {} but the file carries {version}",
begin.version
)));
}
let description = file.header.description.to_string();
self.kvs()
.put_ml_model(&session, &name, &version, &description, file.to_bytes())
.await
.map_err(|err| Status::internal(err.to_string()))?;
Ok(Response::new(checked_response(rpc::ImportMlModelResponse {})?))
}
#[cfg(not(feature = "ml"))]
async fn run_ml_import(
&self,
_frames: tonic::Streaming<rpc::ImportMlModelRequest>,
) -> Result<Response<rpc::ImportMlModelResponse>, Status> {
Err(Status::unimplemented("This server was built without SurrealML support"))
}
async fn start_export(
&self,
session: Session,
config: export::Config,
) -> Result<Export, Status> {
let (namespace, database) =
check_ns_db(&session).map_err(|err| rejected_before_execution(err.to_string()))?;
self.kvs()
.check(
&session,
surrealdb_iam::Action::View,
surrealdb_iam::ResourceKind::Any.on_db(&namespace, &database),
)
.map_err(|err| Status::permission_denied(err.to_string()))?;
let (sender, chunks) = surrealdb::channel::bounded(1);
let task =
self.kvs().export_with_config(&session, sender, config).await.map_err(|err| {
match err.downcast_ref::<TypesError>() {
Some(err) => to_status(err),
None => Status::internal(err.to_string()),
}
})?;
Ok(Export {
chunks,
outcome: tokio::spawn(task),
})
}
#[cfg(feature = "ml")]
async fn start_ml_export(
&self,
session: Session,
name: String,
version: String,
) -> Result<Export, Status> {
let (namespace, database) =
check_ns_db(&session).map_err(|err| rejected_before_execution(err.to_string()))?;
self.kvs()
.check(
&session,
surrealdb_iam::Action::View,
surrealdb_iam::ResourceKind::Model.on_db(&namespace, &database),
)
.map_err(|err| Status::permission_denied(err.to_string()))?;
let info = self
.kvs()
.get_db_model(&namespace, &database, &name, &version)
.await
.map_err(|err| Status::internal(err.to_string()))?
.ok_or_else(|| Status::not_found(format!("Model {name} {version} not found")))?;
let path = format!("ml/{namespace}/{database}/{name}-{version}-{}.surml", info.hash);
let mut data = surrealdb_core::obs::stream(path)
.await
.map_err(|err| Status::internal(format!("Failed to read model file: {err}")))?;
let (sender, chunks) = surrealdb::channel::bounded(1);
let outcome = tokio::spawn(async move {
while let Some(chunk) = data.next().await {
let chunk = chunk.map_err(|err| anyhow::anyhow!("{err}"))?;
if sender.send(chunk.to_vec()).await.is_err() {
break;
}
}
Ok(())
});
Ok(Export {
chunks,
outcome,
})
}
#[cfg(not(feature = "ml"))]
async fn start_ml_export(
&self,
_session: Session,
_name: String,
_version: String,
) -> Result<Export, Status> {
Err(Status::unimplemented("This server was built without SurrealML support"))
}
}
struct Export {
chunks: surrealdb::channel::Receiver<Vec<u8>>,
outcome: tokio::task::JoinHandle<anyhow::Result<()>>,
}
struct SubscriptionGuard {
state: Arc<RpcState>,
live_query_id: Uuid,
subscription_id: Uuid,
owned: bool,
}
impl Drop for SubscriptionGuard {
fn drop(&mut self) {
let state = Arc::clone(&self.state);
let live_query_id = self.live_query_id;
if !self.owned {
if let Some(mut entry) = state.grpc.live_queries.get_mut(&live_query_id)
&& entry.subscriber.as_ref().is_some_and(|s| s.id == self.subscription_id)
{
entry.subscriber = None;
}
return;
}
tokio::spawn(async move {
state.grpc.forget_live_query(&live_query_id);
if let Err(err) = state.grpc.kvs.delete_queries(vec![live_query_id]).await {
error!("Error killing subscription-owned live query {live_query_id}: {err}");
}
});
}
}
enum ByteFrame {
Chunk(rpc::DataChunk),
Trailer(rpc::DataTrailer),
Error(proto::SurrealError),
}
fn shortest(a: Option<Duration>, b: Option<Duration>) -> Option<Duration> {
match (a, b) {
(Some(a), Some(b)) => Some(a.min(b)),
(deadline, None) | (None, deadline) => deadline,
}
}
type QueryStreamRun = Pin<Box<dyn Future<Output = Result<Vec<QueryResult>, TypesError>> + Send>>;
const QUERY_FIRST_BATCH_RECORDS: usize = 16;
struct ReleaseOnDrop {
cancel: CancelHandle,
state: Arc<RpcState>,
session_id: Uuid,
ephemeral: bool,
unregistered: Vec<Uuid>,
}
impl ReleaseOnDrop {
fn new(cancel: CancelHandle, state: Arc<RpcState>, session: &ResolvedSession) -> Self {
Self {
cancel,
state,
session_id: session.id,
ephemeral: session.client.is_none(),
unregistered: Vec::new(),
}
}
fn track_live_query(&mut self, id: Uuid) {
self.unregistered.push(id);
}
fn take_live_queries(&mut self) -> Vec<Uuid> {
std::mem::take(&mut self.unregistered)
}
fn released(&mut self) {
self.ephemeral = false;
}
}
impl Drop for ReleaseOnDrop {
fn drop(&mut self) {
self.cancel.trip();
let unregistered = std::mem::take(&mut self.unregistered);
let ephemeral = self.ephemeral.then_some(self.session_id);
if unregistered.is_empty() && ephemeral.is_none() {
return;
}
let state = Arc::clone(&self.state);
let session_id = self.session_id;
tokio::spawn(async move {
discard_unregistered_live_queries(&state, session_id, unregistered).await;
if let Some(id) = ephemeral {
state.grpc.remove_ephemeral_session(&id).await;
}
});
}
}
struct QueryFraming {
state: Arc<RpcState>,
session: ResolvedSession,
principal: AuthPrincipalSnapshot,
run: Option<QueryStreamRun>,
items: Receiver<QueryStreamItem>,
frames: QueryFrames,
started: Instant,
done: bool,
release: ReleaseOnDrop,
}
struct QueryFrames {
pending: StdHashMap<u32, PendingStatement>,
queued: VecDeque<rpc::QueryResponse>,
batch_records: usize,
value_budget: usize,
terminated: StdHashSet<u32>,
}
#[derive(Default)]
struct PendingStatement {
values: Vec<proto::Value>,
costs: Vec<usize>,
bytes: usize,
batches: u64,
target: usize,
sent: i64,
single: bool,
live_query: Option<Uuid>,
}
impl PendingStatement {
fn new(batch_records: usize) -> Self {
Self {
target: QUERY_FIRST_BATCH_RECORDS.min(batch_records).max(1),
..Default::default()
}
}
}
impl QueryFraming {
fn new(
state: Arc<RpcState>,
session: ResolvedSession,
principal: AuthPrincipalSnapshot,
job: QueryStreamJob,
items: Receiver<QueryStreamItem>,
batch_records: usize,
cancel: CancelHandle,
) -> Self {
let release = ReleaseOnDrop::new(cancel, Arc::clone(&state), &session);
Self {
state,
session,
principal,
run: Some(job.run),
items,
frames: QueryFrames::new(batch_records),
started: Instant::now(),
done: false,
release,
}
}
async fn next(&mut self) -> Option<rpc::QueryResponse> {
loop {
if let Some(frame) = self.frames.pop() {
return Some(frame);
}
if self.done {
if self.session.client.is_none() {
let state = Arc::clone(&self.state);
let id = self.session.id;
self.release.released();
state.grpc.remove_ephemeral_session(&id).await;
}
return None;
}
let run = self.run.as_mut().expect("the execution is driven until it completes");
tokio::select! {
item = self.items.recv() => match item {
Ok(item) => self.absorb(item),
Err(_) => {
let outcome = self.run.take().expect("still running").await;
self.finish(outcome).await;
}
},
outcome = run => {
self.run = None;
while let Ok(item) = self.items.try_recv() {
self.absorb(item);
}
self.finish(outcome).await;
}
}
}
}
fn absorb(&mut self, item: QueryStreamItem) {
if let Some(id) = self.frames.absorb(item) {
self.release.track_live_query(id);
}
}
async fn finish(&mut self, outcome: Result<Vec<QueryResult>, TypesError>) {
self.done = true;
let unregistered = self.release.take_live_queries();
match outcome {
Ok(results) => {
let state = Arc::clone(&self.state);
let session_id = self.session.id;
register_live_queries(&state, session_id, &self.principal, &results).await;
self.frames.end(self.started.elapsed());
}
Err(error) => {
discard_unregistered_live_queries(&self.state, self.session.id, unregistered).await;
self.frames.error(&error);
}
}
}
}
impl QueryFrames {
fn new(batch_records: usize) -> Self {
Self {
pending: StdHashMap::new(),
queued: VecDeque::new(),
batch_records,
value_budget: frame_payload_budget(),
terminated: StdHashSet::new(),
}
}
fn pop(&mut self) -> Option<rpc::QueryResponse> {
self.queued.pop_front()
}
fn end(&mut self, elapsed: Duration) {
self.queued.push_back(rpc::QueryResponse {
frame: Some(rpc::query_response::Frame::End(rpc::QueryEnd {
result_count: self.terminated.len() as u32,
execution_duration: proto::Duration::try_from(elapsed).ok(),
})),
});
}
fn error(&mut self, error: &TypesError) {
self.queued.push_back(rpc::QueryResponse {
frame: Some(rpc::query_response::Frame::Error(to_proto_error(error))),
});
}
fn absorb(&mut self, item: QueryStreamItem) -> Option<Uuid> {
if self.terminated.contains(&(item.index() as u32)) {
return None;
}
match item {
QueryStreamItem::Rows {
index,
values,
} => {
let index = index as u32;
let batch_records = self.batch_records;
self.pending.entry(index).or_insert_with(|| PendingStatement::new(batch_records));
let encoded = match try_values(values.into_iter()) {
Ok(encoded) => encoded,
Err(error) => {
self.pending.remove(&index);
self.queue_terminal(
index,
rpc::QueryStatementKind::Other,
Some(to_proto_error(&error)),
0,
);
return None;
}
};
if let Err(error) = self.accumulate(index, encoded) {
self.pending.remove(&index);
self.queue_terminal(
index,
rpc::QueryStatementKind::Other,
Some(to_proto_error(&error)),
0,
);
}
None
}
QueryStreamItem::Value {
index,
value,
} => {
let index = index as u32;
let batch_records = self.batch_records;
let live_query = match &value {
Value::Uuid(id) => Some(id.into_inner()),
_ => None,
};
match proto::Value::try_from(value).map_err(types_error_from_anyhow).and_then(
|value| match record_cost(&value) {
cost if cost > self.value_budget => {
Err(oversized_record(cost, self.value_budget))
}
cost => Ok((value, cost)),
},
) {
Ok((value, cost)) => {
let entry = self
.pending
.entry(index)
.or_insert_with(|| PendingStatement::new(batch_records));
entry.values.push(value);
entry.costs.push(cost);
entry.bytes += cost;
entry.single = true;
entry.live_query = live_query;
}
Err(error) => {
self.pending.remove(&index);
self.queue_terminal(
index,
rpc::QueryStatementKind::Other,
Some(to_proto_error(&error)),
0,
);
}
}
None
}
QueryStreamItem::Finished {
index,
time,
query_type,
error,
} => {
let index = index as u32;
let kind = match query_type {
QueryType::Live => rpc::QueryStatementKind::Live,
QueryType::Kill => rpc::QueryStatementKind::Kill,
_ => rpc::QueryStatementKind::Other,
};
let live_query = match (query_type, &error) {
(QueryType::Live, None) => {
self.pending.get(&index).and_then(|entry| entry.live_query)
}
_ => None,
};
self.queue_terminal(
index,
kind,
error.map(|e| to_proto_error(&e)),
time.as_nanos(),
);
live_query
}
}
}
fn accumulate(&mut self, index: u32, encoded: Vec<proto::Value>) -> Result<(), TypesError> {
for value in encoded {
let cost = record_cost(&value);
if cost > self.value_budget {
return Err(oversized_record(cost, self.value_budget));
}
let Some(entry) = self.pending.get_mut(&index) else {
return Ok(());
};
if !entry.values.is_empty() && entry.bytes + cost > self.value_budget {
self.flush(index, usize::MAX);
}
let Some(entry) = self.pending.get_mut(&index) else {
return Ok(());
};
entry.values.push(value);
entry.costs.push(cost);
entry.bytes += cost;
let target = entry.target;
if entry.values.len() >= target {
self.flush(index, target);
}
}
Ok(())
}
fn flush(&mut self, index: u32, take: usize) {
let Some(entry) = self.pending.get_mut(&index) else {
return;
};
let take = take.min(entry.values.len());
if take == 0 {
return;
}
let values: Vec<proto::Value> = entry.values.drain(..take).collect();
entry.bytes -= entry.costs.drain(..take).sum::<usize>();
let batch_index = entry.batches;
entry.batches += 1;
entry.sent += values.len() as i64;
entry.target = (entry.target * 2).min(self.batch_records).max(1);
self.queued.push_back(rpc::QueryResponse {
frame: Some(rpc::query_response::Frame::Batch(rpc::QueryBatchFrame {
query_index: index,
batch_index,
statement_kind: rpc::QueryStatementKind::Other as i32,
kind: rpc::QueryResponseKind::Batched as i32,
stats: None,
error: None,
payload: Some(rpc::query_batch_frame::Payload::Values(rpc::ValueBatch {
values,
})),
})),
});
}
fn queue_terminal(
&mut self,
index: u32,
kind: rpc::QueryStatementKind,
error: Option<proto::SurrealError>,
elapsed_nanos: u128,
) {
if !self.terminated.insert(index) {
return;
}
let entry = self.pending.remove(&index).unwrap_or_default();
let (payload, records) = if error.is_some() {
(None, 0)
} else {
let records = entry.sent + entry.values.len() as i64;
(
Some(rpc::query_batch_frame::Payload::Values(rpc::ValueBatch {
values: entry.values,
})),
records,
)
};
let single = error.is_none() && entry.single;
let stats = rpc::QueryStats {
records_returned: records,
bytes_returned: -1,
records_scanned: -1,
bytes_scanned: -1,
execution_duration: u64::try_from(elapsed_nanos)
.ok()
.and_then(|nanos| proto::Duration::try_from(Duration::from_nanos(nanos)).ok()),
};
self.queued.push_back(rpc::QueryResponse {
frame: Some(rpc::query_response::Frame::Batch(rpc::QueryBatchFrame {
query_index: index,
batch_index: entry.batches,
kind: if single {
rpc::QueryResponseKind::Single as i32
} else {
rpc::QueryResponseKind::BatchedFinal as i32
},
statement_kind: kind as i32,
stats: Some(stats),
error,
payload,
})),
});
}
}
async fn register_live_queries(
state: &RpcState,
session_id: Uuid,
principal: &AuthPrincipalSnapshot,
results: &[QueryResult],
) {
if !results.iter().any(|r| r.query_type == QueryType::Live) {
return;
}
let session = live_query_owner(state.grpc.session_map(), session_id, principal).await;
let (namespace, database) = match &session {
Some(owner) => (owner.ns.clone(), owner.db.clone()),
None => {
let orphans = results
.iter()
.filter(|result| matches!(result.query_type, QueryType::Live))
.filter_map(|result| match &result.result {
Ok(Value::Uuid(id)) => Some(id.into_inner()),
_ => None,
})
.collect();
discard_unregistered_live_queries(state, session_id, orphans).await;
return;
}
};
for result in results {
let Ok(Value::Uuid(id)) = &result.result else {
continue;
};
if result.query_type == QueryType::Live {
state.grpc.handle_live(id, session_id, namespace.clone(), database.clone()).await;
}
}
}
async fn discard_unregistered_live_queries(state: &RpcState, session_id: Uuid, ids: Vec<Uuid>) {
if ids.is_empty() {
return;
}
let orphaned = ids.len();
let Err(err) = state.grpc.kvs().delete_queries(ids).await else {
return;
};
error!("Error cleaning up the live queries of an unfinished gRPC streaming query: {err}");
if let Some(observer) = state.grpc.metrics_observer.as_ref() {
let (namespace, database) = match state.grpc.get_session(&session_id).await {
Ok(lock) => {
let session = lock.read().await;
(session.ns.clone(), session.db.clone())
}
Err(_) => (None, None),
};
for _ in 0..orphaned {
observer.record_live_query_orphaned(namespace.as_deref(), database.as_deref());
}
}
}
fn frame_byte_stream<T, F>(export: Export, wrap: F) -> impl Stream<Item = Result<T, Status>> + Send
where
F: Fn(ByteFrame) -> T + Send + 'static,
T: Send + 'static,
{
struct Framing {
export: Option<Export>,
buffer: bytes::BytesMut,
hasher: blake3::Hasher,
streamed: u64,
}
futures::stream::unfold(
(
Framing {
export: Some(export),
buffer: bytes::BytesMut::new(),
hasher: blake3::Hasher::new(),
streamed: 0,
},
wrap,
),
|(mut framing, wrap)| async move {
loop {
if framing.buffer.len() >= EXPORT_CHUNK_SIZE {
let data = framing.buffer.split_to(EXPORT_CHUNK_SIZE).freeze();
let frame = wrap(ByteFrame::Chunk(rpc::DataChunk {
data,
}));
return Some((Ok(frame), (framing, wrap)));
}
let export = framing.export.take()?;
match export.chunks.recv().await {
Ok(chunk) => {
framing.hasher.update(&chunk);
framing.streamed += chunk.len() as u64;
framing.buffer.extend_from_slice(&chunk);
framing.export = Some(export);
}
Err(_) if !framing.buffer.is_empty() => {
let data = std::mem::take(&mut framing.buffer).freeze();
let frame = wrap(ByteFrame::Chunk(rpc::DataChunk {
data,
}));
framing.export = Some(export);
return Some((Ok(frame), (framing, wrap)));
}
Err(_) => {
let frame = match export.outcome.await {
Ok(Ok(())) => ByteFrame::Trailer(rpc::DataTrailer {
bytes: framing.streamed,
blake3: framing.hasher.finalize().to_hex().to_string(),
}),
Ok(Err(err)) => {
error!("gRPC export failed: {err}");
ByteFrame::Error(proto::SurrealError::new(
proto::ErrorKind::Internal,
"The export failed part-way through",
))
}
Err(err) => {
error!("gRPC export task panicked: {err}");
ByteFrame::Error(proto::SurrealError::new(
proto::ErrorKind::Internal,
"The export failed part-way through",
))
}
};
return Some((Ok(wrap(frame)), (framing, wrap)));
}
}
}
},
)
}
fn single_live_query(mut results: Vec<QueryResult>) -> Result<Uuid, Status> {
if results.len() != 1 {
return Err(Status::invalid_argument("Expected exactly one LIVE SELECT statement"));
}
let result = results.remove(0);
if result.query_type != QueryType::Live {
return Err(Status::invalid_argument("Expected a LIVE SELECT statement"));
}
match result.result.map_err(|err| to_status(&err))? {
Value::Uuid(id) => Ok(id.into_inner()),
_ => Err(Status::internal("LIVE SELECT did not return a live query id")),
}
}
fn verify_trailer(
trailer: &rpc::DataTrailer,
streamed: u64,
digest: blake3::Hash,
) -> Result<(), Status> {
if trailer.bytes != streamed {
return Err(Status::data_loss(format!(
"Import trailer declared {} bytes but {streamed} arrived",
trailer.bytes
)));
}
if !trailer.blake3.is_empty() && !trailer.blake3.eq_ignore_ascii_case(&digest.to_hex()) {
return Err(Status::data_loss("Import trailer checksum did not match the streamed bytes"));
}
Ok(())
}
fn try_values(values: impl Iterator<Item = Value>) -> Result<Vec<proto::Value>, TypesError> {
values.map(|value| proto::Value::try_from(value).map_err(types_error_from_anyhow)).collect()
}
fn batch_records(requested: u32) -> usize {
match requested as usize {
0 => QUERY_BATCH_RECORDS,
n => n.min(QUERY_BATCH_RECORDS),
}
}
fn to_proto_notification(
notification: &Notification,
action: surrealdb_types::Action,
) -> Result<rpc::Notification, TypesError> {
let action = match action {
surrealdb_types::Action::Create => rpc::Action::Created,
surrealdb_types::Action::Update => rpc::Action::Updated,
surrealdb_types::Action::Delete => rpc::Action::Deleted,
surrealdb_types::Action::Killed | surrealdb_types::Action::Error => {
rpc::Action::Unspecified
}
};
let record_id = match ¬ification.record {
Value::RecordId(record) => {
Some(proto::RecordId::try_from(record.clone()).map_err(types_error_from_anyhow)?)
}
_ => None,
};
Ok(rpc::Notification {
live_query_id: Some(proto::Uuid::from_uuid(notification.id.into_inner())),
action: action as i32,
record_id,
value: Some(
proto::Value::try_from(notification.result.clone()).map_err(types_error_from_anyhow)?,
),
cursor: None,
})
}
fn to_tokens(result: DbResult) -> Result<rpc::Tokens, Status> {
let DbResult::Other(value) = result else {
return Err(Status::internal("Authentication did not return a token"));
};
if matches!(value, Value::None | Value::Null) {
return Ok(rpc::Tokens::default());
}
let (access, refresh) = match Token::from_value(value) {
Ok(Token::Access(access)) => (access, String::new()),
Ok(Token::WithRefresh {
access,
refresh,
}) => (access, refresh),
Err(err) => {
return Err(Status::internal(format!("Authentication returned no token: {err}")));
}
};
Ok(rpc::Tokens {
access,
refresh,
expires_at: None,
refresh_expires_at: None,
})
}
fn token_value(access: String, refresh: String) -> Value {
let token = if refresh.is_empty() {
Token::Access(access)
} else {
Token::WithRefresh {
access,
refresh,
}
};
token.into_value()
}
fn access_credentials(access: rpc::AccessMethod) -> Result<surrealdb_types::Object, Status> {
let method =
access.method.ok_or_else(|| Status::invalid_argument("Expected an access method"))?;
let mut object = surrealdb_types::Object::new();
let mut set = |key: &str, value: String| {
if !value.is_empty() {
object.insert(key.to_string(), Value::String(value));
}
};
match method {
rpc::access_method::Method::User(user) => {
set("ns", user.namespace);
set("db", user.database);
set("user", user.username);
set("pass", user.password);
set("ac", user.access);
}
rpc::access_method::Method::Bearer(bearer) => {
set("ns", bearer.namespace);
set("db", bearer.database);
set("ac", bearer.access);
set("key", bearer.key);
}
rpc::access_method::Method::Record(record) => {
return record_credentials(record);
}
}
Ok(object)
}
fn record_credentials(record: rpc::RecordCredentials) -> Result<surrealdb_types::Object, Status> {
let mut object = match record.variables {
Some(variables) => match from_proto_variables(variables)? {
Value::Object(object) => object,
_ => surrealdb_types::Object::new(),
},
None => surrealdb_types::Object::new(),
};
let mut set = |key: &str, value: String| {
if !value.is_empty() {
object.insert(key.to_string(), Value::String(value));
}
};
set("ns", record.namespace);
set("db", record.database);
set("ac", record.access);
Ok(object)
}
fn from_proto_export_config(config: rpc::ExportConfig) -> export::Config {
use export::TableConfig;
let tables = match config.tables.and_then(|tables| tables.selection) {
Some(rpc::export_config::tables::Selection::All(_)) | None => TableConfig::All,
Some(rpc::export_config::tables::Selection::None(_)) => TableConfig::None,
Some(rpc::export_config::tables::Selection::Selected(selected)) => {
TableConfig::Some(selected.tables)
}
Some(rpc::export_config::tables::Selection::Excluded(excluded)) => {
TableConfig::Exclude(export::ExcludedTables {
exclude: excluded.tables,
})
}
};
export::Config {
database_definition: export::Config::default().database_definition,
users: config.users,
accesses: config.accesses,
params: config.params,
functions: config.functions,
analyzers: config.analyzers,
tables,
versions: config.versions,
records: config.records,
sequences: config.sequences,
apis: config.apis,
buckets: config.buckets,
modules: config.modules,
configs: config.configs,
}
}
fn from_nullable(value: Option<rpc::NullableString>) -> Value {
match value.and_then(|value| value.value) {
Some(rpc::nullable_string::Value::Some(value)) => Value::String(value),
Some(rpc::nullable_string::Value::Null(_)) => Value::Null,
None => Value::None,
}
}
fn to_proto_value(value: Value) -> Result<proto::Value, Status> {
proto::Value::try_from(value).map_err(|err| Status::internal(err.to_string()))
}
fn from_proto_value(value: proto::Value) -> Result<Value, Status> {
Value::try_from(value).map_err(|err| Status::invalid_argument(err.to_string()))
}
fn from_proto_variables(variables: proto::Variables) -> Result<Value, Status> {
let object = proto::Object {
items: variables.variables,
};
from_proto_value(proto::Value {
value: Some(proto::value::Value::Object(object)),
})
}
fn to_uuid(uuid: &proto::Uuid) -> Result<Uuid, Status> {
uuid.to_uuid().map_err(|err| Status::invalid_argument(err.to_string()))
}
fn to_proto_details(details: &TypesErrorDetails) -> Option<(String, Option<Value>)> {
let Value::Object(mut outer) = details.clone().into_value() else {
return None;
};
let Some(Value::Object(mut reason)) = outer.remove("details") else {
return None;
};
let Some(Value::String(kind)) = reason.remove("kind") else {
return None;
};
Some((kind, reason.remove("details")))
}
pub(crate) const REJECTED_BEFORE_EXECUTION: &str = "surreal-rejected-before-execution";
fn rejected_before_execution(message: impl Into<String>) -> Status {
let mut status = Status::failed_precondition(message.into());
status.metadata_mut().insert(REJECTED_BEFORE_EXECUTION, "1".parse().expect("a literal 1"));
status
}
const TRUNCATION_MARKER: &str = " [...truncated]";
fn to_proto_error(error: &TypesError) -> proto::SurrealError {
bounded_error(proto_error(error), frame_payload_budget())
}
fn bounded_error(mut error: proto::SurrealError, budget: usize) -> proto::SurrealError {
use surrealdb_types::prost::Message;
if error.encoded_len() <= budget {
return error;
}
error.cause = None;
if error.encoded_len() <= budget {
return error;
}
error.details = None;
if error.encoded_len() <= budget {
return error;
}
let mut fixed = error.clone();
fixed.message = String::new();
let room =
budget.saturating_sub(fixed.encoded_len() + WIRE_FIELD_OVERHEAD + TRUNCATION_MARKER.len());
error.message = format!("{}{TRUNCATION_MARKER}", truncate_on_boundary(&error.message, room));
error
}
fn truncate_on_boundary(text: &str, bytes: usize) -> &str {
if text.len() <= bytes {
return text;
}
let mut end = bytes;
while end > 0 && !text.is_char_boundary(end) {
end -= 1;
}
&text[..end]
}
fn proto_error(error: &TypesError) -> proto::SurrealError {
let kind = if error.is_validation() {
proto::ErrorKind::Validation
} else if error.is_configuration() {
proto::ErrorKind::Configuration
} else if error.is_query() {
proto::ErrorKind::Query
} else if error.is_serialization() {
proto::ErrorKind::Serialization
} else if error.is_not_allowed() {
proto::ErrorKind::NotAllowed
} else if error.is_not_found() {
proto::ErrorKind::NotFound
} else if error.is_already_exists() {
proto::ErrorKind::AlreadyExists
} else if error.is_connection() {
proto::ErrorKind::Connection
} else if error.is_thrown() {
proto::ErrorKind::Thrown
} else if error.is_context() {
proto::ErrorKind::Context
} else {
proto::ErrorKind::Internal
};
let mut out = proto::SurrealError::new(kind, error.message());
if let Some((reason, content)) = to_proto_details(error.details()) {
out =
out.with_details(reason, content.and_then(|value| proto::Value::try_from(value).ok()));
}
if matches!(error.query_details(), Some(QueryError::TransactionConflict)) {
out = out.with_retry(None);
}
if let Some(cause) = error.cause() {
out = out.with_cause(proto_error(cause));
}
out
}
const STATUS_MESSAGE_BYTES: usize = 2 << 10;
fn bounded_status_message(message: &str) -> String {
if message.len() <= STATUS_MESSAGE_BYTES {
return message.to_string();
}
let room = STATUS_MESSAGE_BYTES - TRUNCATION_MARKER.len();
format!("{}{TRUNCATION_MARKER}", truncate_on_boundary(message, room))
}
fn to_status(error: &TypesError) -> Status {
let message = bounded_status_message(error.message());
if error.is_validation() {
Status::invalid_argument(message)
} else if error.is_configuration() {
Status::unimplemented(message)
} else if error.is_query() || error.is_thrown() {
Status::aborted(message)
} else if error.is_not_allowed() {
Status::permission_denied(message)
} else if error.is_not_found() {
Status::not_found(message)
} else if error.is_already_exists() {
Status::already_exists(message)
} else if error.is_connection() {
Status::unavailable(message)
} else {
Status::internal(message)
}
}
fn import_failure_status(results: &[QueryResult]) -> Option<Status> {
const FIRST_FAILURE_BUDGET: usize = 512;
let mut failures = results.iter().filter_map(|r| r.result.as_ref().err());
let first = failures.next()?;
let further = failures.count();
let first = first.to_string();
let mut message = truncate_on_boundary(&first, FIRST_FAILURE_BUDGET).to_owned();
if message.len() < first.len() {
message.push_str(TRUNCATION_MARKER);
}
if further > 0 {
message.push_str(&format!(" (and {further} further statements failed)"));
}
message.push_str(
"; the import is not transactional, so the statements before this one have been applied",
);
Some(Status::failed_precondition(message))
}
#[cfg(test)]
mod tests {
use surrealdb_rpc::capabilities::Capabilities;
use surrealdb_types::{AuthError, NotAllowedError, Object, SerializationError};
use super::*;
async fn service() -> GrpcService {
service_for(Session::default()).await
}
async fn service_for(caller: Session) -> GrpcService {
let datastore = Datastore::builder()
.with_capabilities(Capabilities::all())
.build_with_path("memory")
.await
.expect("datastore");
GrpcService::new(Arc::new(RpcState::new(datastore)), caller)
}
async fn register_live_query(service: &GrpcService, session_id: Uuid) -> Uuid {
if !service.rpc().sessions.contains_key(&session_id) {
service.rpc().attach(session_id).await.expect("attach");
}
let id = Uuid::new_v4();
service.rpc().handle_live(&id, session_id, None, None).await;
id
}
fn registered_live_queries(service: &GrpcService, session_id: Uuid) -> Vec<Uuid> {
let mut ids: Vec<Uuid> = service
.rpc()
.live_queries
.iter()
.filter(|entry| entry.value().session_id == session_id)
.map(|entry| *entry.key())
.collect();
ids.sort();
ids
}
fn live_query_count(service: &GrpcService, session_id: Uuid) -> Option<usize> {
service.rpc().live_query_counts.get(&session_id).map(|entry| *entry.value())
}
async fn begin_txn(service: &GrpcService, session_id: Uuid) -> Uuid {
let DbResult::Other(Value::Uuid(txn)) =
service.rpc().begin(None, session_id).await.expect("begin")
else {
panic!("begin should answer with a transaction id");
};
txn.into_inner()
}
#[tokio::test]
async fn a_credentialed_caller_may_use_the_session_it_attached() {
let service = service_for(Session::owner()).await;
let session_id = Uuid::new_v4();
service.rpc().attach(session_id).await.expect("attach");
service
.rpc()
.verify_caller_for_session(&session_id, service.caller.au.as_ref())
.await
.expect("the caller that attached the session may use it");
}
#[tokio::test]
async fn a_credentialed_caller_may_not_use_another_principals_session() {
let service = service_for(Session::owner()).await;
let session_id = Uuid::new_v4();
service.rpc().attach(session_id).await.expect("attach");
{
let session = service.rpc().get_session(&session_id).await.expect("session");
let mut session = session.write().await;
session.au = Arc::new(Auth::for_ns(surrealdb_iam::Role::Owner, "test-ns"));
}
let refused = service
.rpc()
.verify_caller_for_session(&session_id, service.caller.au.as_ref())
.await
.expect_err("a different principal must be refused");
assert!(refused.is_not_found(), "expected a not-found error, got {refused:?}");
}
#[tokio::test]
async fn a_transaction_is_only_usable_by_the_session_that_opened_it() {
let service = service().await;
let owner = Uuid::new_v4();
service.rpc().attach(owner).await.expect("attach");
let DbResult::Other(Value::Uuid(txn)) =
service.rpc().begin(None, owner).await.expect("begin")
else {
panic!("begin should answer with a transaction id");
};
let txn = txn.into_inner();
assert!(service.rpc().transaction_belongs_to(&txn, owner));
assert!(!service.rpc().transaction_belongs_to(&txn, Uuid::new_v4()));
assert!(!service.rpc().transaction_belongs_to(&Uuid::new_v4(), owner));
}
#[tokio::test]
async fn a_transaction_needs_an_attached_session() {
let service = service().await;
let ephemeral = Uuid::new_v4();
service
.rpc()
.register_ephemeral_session(ephemeral, Arc::new(RwLock::new(Session::default())));
let refused = service
.rpc()
.begin(None, ephemeral)
.await
.expect_err("an ephemeral session must not open a transaction");
assert!(refused.is_validation(), "expected a validation error, got {refused:?}");
}
#[tokio::test]
async fn a_subscriber_that_stops_reading_is_ended() {
let service = service().await;
let owner = Uuid::new_v4();
let live_query_id = register_live_query(&service, owner).await;
let subscription =
service.attach_subscription(owner, live_query_id, false).expect("subscribe");
let _stream = subscription.into_inner();
let notification = Notification::new(
live_query_id.into(),
None,
surrealdb_types::Action::Create,
Value::None,
Value::None,
);
for _ in 0..*GRPC_NOTIFICATION_BUFFER {
assert!(service.rpc().dispatch_notification(¬ification).await);
assert!(
service
.rpc()
.live_queries
.get(&live_query_id)
.is_some_and(|lq| lq.subscriber.is_some()),
"the subscription must survive while its buffer has room"
);
}
assert!(service.rpc().dispatch_notification(¬ification).await);
assert!(
service
.rpc()
.live_queries
.get(&live_query_id)
.is_some_and(|lq| lq.subscriber.is_none()),
"a full buffer must end the subscription rather than queue behind it"
);
}
#[tokio::test]
async fn shutdown_cancels_the_transactions_clients_left_open() {
let service = service().await;
let session_id = Uuid::new_v4();
service.rpc().attach(session_id).await.expect("attach");
service.rpc().begin(None, session_id).await.expect("begin");
assert_eq!(service.rpc().transactions.len(), 1);
service.rpc().cleanup_all_txns().await;
assert!(service.rpc().transactions.is_empty(), "shutdown must cancel every transaction");
}
#[tokio::test]
async fn a_sessions_live_queries_are_reclaimed_and_no_others() {
let service = service().await;
let session_id = Uuid::new_v4();
let first = register_live_query(&service, session_id).await;
let second = register_live_query(&service, session_id).await;
let bystander_session = Uuid::new_v4();
let bystander = register_live_query(&service, bystander_session).await;
let mut registered = vec![first, second];
registered.sort();
assert_eq!(
registered_live_queries(&service, session_id),
registered,
"both live queries are registered under the session that made them"
);
assert_eq!(
live_query_count(&service, session_id),
Some(2),
"and both are counted, so the session's teardown searches for them"
);
service.rpc().cleanup_lqs(&session_id).await;
assert!(!service.rpc().live_queries.contains_key(&first), "the first is unregistered");
assert!(!service.rpc().live_queries.contains_key(&second), "and so is the second");
assert_eq!(
live_query_count(&service, session_id),
None,
"the count entry is released rather than left behind at zero"
);
assert!(
service.rpc().live_queries.contains_key(&bystander),
"another session's live query survives"
);
assert_eq!(
live_query_count(&service, bystander_session),
Some(1),
"and stays counted under its own session"
);
}
#[tokio::test]
async fn a_registered_live_query_is_counted_and_an_unregistered_one_is_not() {
let service = service().await;
let session_id = Uuid::new_v4();
assert_eq!(
live_query_count(&service, session_id),
None,
"a session that has registered nothing is not counted"
);
let lqid = register_live_query(&service, session_id).await;
assert_eq!(
live_query_count(&service, session_id),
Some(1),
"registering counts against the session"
);
service.rpc().forget_live_query(&lqid);
assert_eq!(
live_query_count(&service, session_id),
None,
"and unregistering releases the count"
);
}
#[tokio::test]
async fn a_live_query_whose_session_went_away_is_not_left_registered() {
let service = service().await;
let session_id = Uuid::new_v4();
let lqid = Uuid::new_v4();
service.rpc().handle_live(&lqid, session_id, None, None).await;
assert!(
!service.rpc().live_queries.contains_key(&lqid),
"the registration must not survive the session it belongs to"
);
assert_eq!(live_query_count(&service, session_id), None, "and must leave no count behind");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_registration_and_removal_keep_the_count_covering_the_registry() {
let service = service().await;
for _ in 0..2000 {
let session_id = Uuid::new_v4();
service.rpc().attach(session_id).await.expect("attach");
let lqid = Uuid::new_v4();
let g1 = Arc::clone(&service.state.grpc);
let g2 = Arc::clone(&service.state.grpc);
let register = tokio::spawn(async move {
g1.handle_live(&lqid, session_id, None, None).await;
});
let end = tokio::spawn(async move {
g2.forget_live_query(&lqid);
});
let _ = tokio::join!(register, end);
let registered = registered_live_queries(&service, session_id);
assert!(
registered.is_empty() || live_query_count(&service, session_id).is_some(),
"the registry holds {registered:?} for a session with no count"
);
service.rpc().cleanup_lqs(&session_id).await;
}
}
#[tokio::test]
async fn registering_a_live_query_racing_detach_strands_nothing() {
let service = service().await;
for _ in 0..100 {
let session_id = Uuid::new_v4();
service.rpc().attach(session_id).await.expect("attach a fresh session");
let lqid = Uuid::new_v4();
let g1 = Arc::clone(&service.state.grpc);
let g2 = Arc::clone(&service.state.grpc);
let register = tokio::spawn(async move {
g1.handle_live(&lqid, session_id, None, None).await;
});
let detach = tokio::spawn(async move {
let _ = g2.del_session(&session_id).await;
});
let _ = tokio::join!(register, detach);
assert!(
!service.rpc().live_queries.contains_key(&lqid),
"live query {lqid} was stranded under detached session"
);
assert_eq!(
live_query_count(&service, session_id),
None,
"a count was stranded under detached session"
);
}
}
#[tokio::test]
async fn a_sessions_transactions_are_reclaimed_and_no_others() {
let service = service().await;
let session_id = Uuid::new_v4();
service.rpc().attach(session_id).await.expect("attach");
let bystander_session = Uuid::new_v4();
service.rpc().attach(bystander_session).await.expect("attach");
let txn = begin_txn(&service, session_id).await;
let bystander = begin_txn(&service, bystander_session).await;
assert!(
service.rpc().transaction_counts.contains_key(&session_id),
"an open transaction holds a slot, so the session's teardown searches for it"
);
service.rpc().cleanup_txns(&session_id).await;
assert!(!service.rpc().transactions.contains_key(&txn), "the transaction is cancelled");
assert!(
!service.rpc().transaction_counts.contains_key(&session_id),
"and its slot released"
);
assert!(
service.rpc().transactions.contains_key(&bystander),
"another session's transaction survives"
);
assert!(
service.rpc().transaction_counts.contains_key(&bystander_session),
"and keeps holding its own slot"
);
}
#[tokio::test]
async fn finishing_a_transaction_releases_the_sessions_slot() {
for (name, finish) in [("commit", true), ("cancel", false)] {
let service = service().await;
let session_id = Uuid::new_v4();
service.rpc().attach(session_id).await.expect("attach");
let txn = begin_txn(&service, session_id).await;
let outcome = match finish {
true => service.rpc().commit(Some(txn), session_id, Array::new()).await,
false => service.rpc().cancel(Some(txn), session_id, Array::new()).await,
};
outcome.unwrap_or_else(|err| panic!("{name} should succeed: {err:?}"));
assert!(
!service.rpc().transaction_counts.contains_key(&session_id),
"{name} must leave no slot reserved"
);
assert!(
!service.rpc().transactions.contains_key(&txn),
"{name} must remove the transaction"
);
}
}
#[tokio::test]
async fn reset_reclaims_what_the_session_holds_and_keeps_the_session() {
let service = service().await;
let session_id = Uuid::new_v4();
service.rpc().attach(session_id).await.expect("attach");
let txn = begin_txn(&service, session_id).await;
let lqid = register_live_query(&service, session_id).await;
service.rpc().reset(session_id).await.expect("reset");
assert!(
!service.rpc().transactions.contains_key(&txn),
"reset must cancel the session's open transaction"
);
assert!(
!service.rpc().live_queries.contains_key(&lqid),
"reset must end the session's live query"
);
assert!(
service.rpc().sessions.contains_key(&session_id),
"and must leave the session itself attached"
);
}
#[tokio::test]
async fn a_live_query_is_only_reachable_by_its_own_session() {
let service = service().await;
let owner = Uuid::new_v4();
let live_query_id = register_live_query(&service, owner).await;
let refused = service
.attach_subscription(Uuid::new_v4(), live_query_id, false)
.err()
.expect("another session must be refused");
assert_eq!(refused.code(), tonic::Code::NotFound);
let Ok(_subscription) = service.attach_subscription(owner, live_query_id, false) else {
panic!("the owner may subscribe");
};
}
#[tokio::test]
async fn a_second_subscriber_is_refused() {
let service = service().await;
let owner = Uuid::new_v4();
let live_query_id = register_live_query(&service, owner).await;
let Ok(_subscription) = service.attach_subscription(owner, live_query_id, false) else {
panic!("the first subscriber should be accepted");
};
let refused = service
.attach_subscription(owner, live_query_id, false)
.err()
.expect("the second subscriber must be refused");
assert_eq!(refused.code(), tonic::Code::AlreadyExists);
}
#[tokio::test]
async fn dropping_a_subscription_releases_it() {
let service = service().await;
let owner = Uuid::new_v4();
let live_query_id = register_live_query(&service, owner).await;
let Ok(subscription) = service.attach_subscription(owner, live_query_id, false) else {
panic!("the first subscriber should be accepted");
};
drop(subscription);
let Ok(_resubscribed) = service.attach_subscription(owner, live_query_id, false) else {
panic!("re-subscribing after the first stream was dropped");
};
}
#[tokio::test]
async fn an_unknown_live_query_is_refused() {
let service = service().await;
let refused = service
.attach_subscription(Uuid::new_v4(), Uuid::new_v4(), false)
.err()
.expect("an unknown live query must be refused");
assert_eq!(refused.code(), tonic::Code::NotFound);
}
async fn realtime_session(service: &GrpcService) -> Uuid {
let id = Uuid::new_v4();
service.rpc().attach(id).await.expect("attach");
let lock = service.rpc().get_session(&id).await.expect("the attached session");
let mut session = lock.write().await;
*session = Session::owner().with_ns("test").with_db("test").with_rt(true);
session.id = Some(id);
drop(session);
id
}
fn context(session_id: Uuid) -> rpc::RequestContext {
rpc::RequestContext {
session: Some(proto::Uuid::from_uuid(session_id)),
transaction: None,
timeout: None,
}
}
async fn write(service: &GrpcService, statement: &str) {
let session = Session::owner().with_ns("test").with_db("test");
let results =
service.kvs().execute(statement, &session, None).await.expect("the write runs");
for result in results {
result.result.expect("the write succeeds");
}
}
async fn define_watched_table(service: &GrpcService) {
write(service, "DEFINE NAMESPACE test; DEFINE DATABASE test; DEFINE TABLE thing").await;
}
async fn commit_live_query(service: &GrpcService) -> Uuid {
let session = Session::owner().with_ns("test").with_db("test").with_rt(true);
let mut results = service
.kvs()
.execute("LIVE SELECT * FROM thing", &session, None)
.await
.expect("the live query runs");
match results.remove(0).result.expect("the live query succeeds") {
Value::Uuid(id) => id.into_inner(),
value => panic!("a live query answers with its id, got {value:?}"),
}
}
async fn live_query_rows(service: &GrpcService) -> usize {
let session = Session::owner().with_ns("test").with_db("test");
let mut results = service
.kvs()
.execute("INFO FOR TABLE thing", &session, None)
.await
.expect("the table info runs");
let info = match results.remove(0).result.expect("the table info succeeds") {
Value::Object(info) => info,
value => panic!("table info answers with an object, got {value:?}"),
};
match info.get("lives") {
Some(Value::Object(lives)) => lives.len(),
other => panic!("table info reports its live queries, got {other:?}"),
}
}
async fn read_to_live_statement(frames: &mut ResponseStream<rpc::QueryResponse>) {
loop {
let frame = frames.next().await.expect("a frame").expect("a frame");
if let Some(rpc::query_response::Frame::Batch(batch)) = &frame.frame
&& batch.statement_kind == rpc::QueryStatementKind::Live as i32
{
return;
}
}
}
#[tokio::test]
async fn an_abandoned_stream_deletes_the_live_query_it_committed() {
let service = service().await;
let session_id = realtime_session(&service).await;
let context = context(session_id);
define_watched_table(&service).await;
let response = service
.stream_query(
Some(&context),
"LIVE SELECT * FROM thing; SLEEP 30s".to_string(),
None,
QUERY_BATCH_RECORDS,
)
.await
.expect("the stream starts");
let mut frames = response.into_inner();
read_to_live_statement(&mut frames).await;
assert_eq!(
live_query_rows(&service).await,
1,
"the live query is committed before the stream is abandoned",
);
drop(frames);
let mut deleted = false;
for _ in 0..100 {
if live_query_rows(&service).await == 0 {
deleted = true;
break;
}
tokio::time::sleep(Duration::from_millis(20)).await;
}
assert!(deleted, "an abandoned stream's live query must not outlive it");
assert!(
service.rpc().live_queries.is_empty(),
"an abandoned stream registers nothing either",
);
}
#[tokio::test]
async fn a_failed_run_deletes_the_live_queries_it_committed() {
let service = service().await;
let session_id = realtime_session(&service).await;
define_watched_table(&service).await;
let id = commit_live_query(&service).await;
assert_eq!(
live_query_rows(&service).await,
1,
"the live query is committed before the run fails",
);
let (items, buffered) = bounded(QUERY_STREAM_BUFFER);
for item in surrealdb_rpc::items_for_result(
0,
QueryResult {
time: Duration::ZERO,
result: Ok(Value::Uuid(surrealdb_types::Uuid::from(id))),
query_type: QueryType::Live,
},
) {
items.send(item).await.expect("the items fit the stream's buffer");
}
drop(items);
let job = QueryStreamJob {
statement_count: 1,
run: Box::pin(async {
Err(TypesError::internal("the run failed as a whole".to_string()))
}),
};
let mut framing = QueryFraming::new(
Arc::clone(&service.state),
ResolvedSession {
id: session_id,
client: Some(session_id),
},
AuthPrincipalSnapshot::capture(&Session::default()),
job,
buffered,
QUERY_BATCH_RECORDS,
CancelHandle::new(),
);
let mut failed = false;
while let Some(frame) = framing.next().await {
failed |= matches!(frame.frame, Some(rpc::query_response::Frame::Error(_)));
}
assert!(failed, "the run's failure ends the stream");
assert_eq!(
live_query_rows(&service).await,
0,
"a failed run's live query must not outlive it",
);
assert!(service.rpc().live_queries.is_empty(), "a failed run registers nothing either");
}
#[tokio::test]
async fn a_detached_sessions_live_query_is_deleted_not_registered() {
let service = service().await;
let session_id = realtime_session(&service).await;
let context = context(session_id);
define_watched_table(&service).await;
let response = service
.stream_query(
Some(&context),
"LIVE SELECT * FROM thing; SLEEP 200ms".to_string(),
None,
QUERY_BATCH_RECORDS,
)
.await
.expect("the stream starts");
let mut frames = response.into_inner();
read_to_live_statement(&mut frames).await;
service.rpc().detach(session_id).await.expect("detach the session");
while let Some(frame) = frames.next().await {
frame.expect("a frame");
}
assert!(
service.rpc().live_queries.is_empty(),
"a detached session's live query must not be registered",
);
assert_eq!(
live_query_rows(&service).await,
0,
"a detached session's live query must not be left in the datastore",
);
}
#[test]
fn framing_reports_the_live_queries_it_frames() {
let id = surrealdb_types::Uuid::new_v4();
let mut framing = QueryFrames::new(QUERY_BATCH_RECORDS);
let live = QueryResult {
time: Duration::ZERO,
result: Ok(Value::Uuid(id)),
query_type: QueryType::Live,
};
let reported: Vec<Uuid> = surrealdb_rpc::items_for_result(0, live)
.into_iter()
.filter_map(|item| framing.absorb(item))
.collect();
assert_eq!(reported, vec![id.into_inner()], "the live query is reported once, on finish");
let mut framing = QueryFrames::new(QUERY_BATCH_RECORDS);
let failed = QueryResult {
time: Duration::ZERO,
result: Err(TypesError::internal("no".to_string())),
query_type: QueryType::Live,
};
let other = QueryResult {
time: Duration::ZERO,
result: Ok(Value::Uuid(surrealdb_types::Uuid::new_v4())),
query_type: QueryType::Other,
};
let reported: Vec<Uuid> = [failed, other]
.into_iter()
.enumerate()
.flat_map(|(index, result)| surrealdb_rpc::items_for_result(index, result))
.filter_map(|item| framing.absorb(item))
.collect();
assert!(reported.is_empty(), "nothing else is taken for a live query: {reported:?}");
}
fn frames(results: Vec<QueryResult>) -> Vec<rpc::QueryResponse> {
frames_with(results, QUERY_BATCH_RECORDS)
}
fn frames_with(results: Vec<QueryResult>, batch_records: usize) -> Vec<rpc::QueryResponse> {
let mut framing = QueryFrames::new(batch_records);
let statement_count = results.len() as u32;
let mut frames = vec![rpc::QueryResponse {
frame: Some(rpc::query_response::Frame::Begin(rpc::QueryBegin {
query_id: Some(proto::Uuid::from_uuid(Uuid::new_v4())),
statement_count,
})),
}];
for (index, result) in results.into_iter().enumerate() {
for item in surrealdb_rpc::items_for_result(index, result) {
framing.absorb(item);
}
}
framing.end(Duration::from_millis(3));
while let Some(frame) = framing.pop() {
frames.push(frame);
}
frames
}
fn batch(response: &rpc::QueryResponse) -> &rpc::QueryBatchFrame {
match response.frame.as_ref() {
Some(rpc::query_response::Frame::Batch(batch)) => batch,
_ => panic!("expected a batch frame"),
}
}
fn result(value: Value) -> QueryResult {
QueryResult {
time: Duration::from_millis(1),
result: Ok(value),
query_type: QueryType::Other,
}
}
#[test]
fn a_terminated_statement_takes_no_further_rows() {
let mut frames = QueryFrames::new(QUERY_BATCH_RECORDS);
frames.queue_terminal(
0,
rpc::QueryStatementKind::Other,
Some(to_proto_error(&TypesError::internal("unencodable".to_string()))),
0,
);
for _ in 0..4 {
frames.absorb(QueryStreamItem::Rows {
index: 0,
values: vec![Value::Bool(true); QUERY_BATCH_RECORDS],
});
}
frames.absorb(QueryStreamItem::Finished {
index: 0,
time: Duration::ZERO,
query_type: QueryType::Other,
error: None,
});
frames.end(Duration::ZERO);
let mut produced = Vec::new();
while let Some(frame) = frames.pop() {
produced.push(frame);
}
assert_eq!(produced.len(), 2, "one terminal batch and one end frame, nothing after");
assert!(
matches!(produced[0].frame, Some(rpc::query_response::Frame::Batch(_))),
"the terminal batch"
);
assert!(matches!(produced[1].frame, Some(rpc::query_response::Frame::End(_))));
}
#[test]
fn a_streaming_query_takes_the_sooner_deadline() {
let short = Duration::from_secs(1);
let long = Duration::from_secs(60);
assert_eq!(shortest(Some(short), Some(long)), Some(short));
assert_eq!(shortest(Some(long), Some(short)), Some(short));
assert_eq!(shortest(Some(short), None), Some(short), "the client's alone still applies");
assert_eq!(shortest(None, Some(long)), Some(long), "the server's alone still applies");
assert_eq!(shortest(None, None), None, "neither configured means no deadline");
}
#[test]
fn a_statement_is_terminated_only_once() {
let mut frames = QueryFrames::new(QUERY_BATCH_RECORDS);
frames.queue_terminal(0, rpc::QueryStatementKind::Other, None, 0);
frames.absorb(QueryStreamItem::Finished {
index: 0,
time: Duration::ZERO,
query_type: QueryType::Other,
error: None,
});
frames.end(Duration::ZERO);
let mut produced = Vec::new();
while let Some(frame) = frames.pop() {
produced.push(frame);
}
assert_eq!(produced.len(), 2, "one terminal batch and one end frame");
assert!(matches!(produced[0].frame, Some(rpc::query_response::Frame::Batch(_))));
match produced[1].frame.as_ref() {
Some(rpc::query_response::Frame::End(end)) => {
assert_eq!(end.result_count, 1, "the statement is counted once");
}
other => panic!("expected an end frame, got {other:?}"),
}
}
#[test]
fn query_frames_are_bracketed_by_begin_and_end() {
let frames = frames(vec![result(Value::None), result(Value::None)]);
assert_eq!(frames.len(), 4);
match frames[0].frame.as_ref() {
Some(rpc::query_response::Frame::Begin(begin)) => {
assert_eq!(begin.statement_count, 2);
}
_ => panic!("expected a begin frame"),
}
match frames[3].frame.as_ref() {
Some(rpc::query_response::Frame::End(end)) => {
assert_eq!(end.result_count, 2, "every parsed statement produced a result");
assert!(
end.execution_duration.is_some(),
"the whole-query duration is what a client cannot sum from the per-statement ones"
);
}
_ => panic!("expected an end frame"),
}
}
#[test]
fn a_client_batch_size_is_honoured_and_clamped() {
assert_eq!(batch_records(0), QUERY_BATCH_RECORDS, "zero means the server chooses");
assert_eq!(batch_records(1), 1, "a client may ask for one record per batch");
assert_eq!(batch_records(u32::MAX), QUERY_BATCH_RECORDS, "clamped, not refused");
let rows = Array::from(
(0..3).map(|i| Value::Number(surrealdb_types::Number::Int(i))).collect::<Vec<_>>(),
);
let frames = frames_with(vec![result(Value::Array(rows))], batch_records(1));
let sizes = batch_sizes(&frames);
assert!(
sizes.iter().all(|n| *n <= 1),
"a client asking for one record per batch gets no more than one, {sizes:?}"
);
assert_eq!(sizes.iter().sum::<usize>(), 3, "every record is still sent once");
}
fn batch_sizes(frames: &[rpc::QueryResponse]) -> Vec<usize> {
frames[1..frames.len() - 1]
.iter()
.map(|response| match batch(response).payload.as_ref() {
Some(rpc::query_batch_frame::Payload::Values(values)) => values.values.len(),
None => 0,
_ => panic!("expected a value batch"),
})
.collect()
}
#[test]
fn the_transport_reserve_clears_worst_case_compression() {
for limit in [
crate::cnf::GRPC_MIN_MESSAGE_SIZE,
4 << 20,
128 << 20,
crate::cnf::GRPC_MAX_MESSAGE_CEILING,
] {
let payload = limit - transport_reserve(limit);
assert!(
payload + payload / 2048 <= limit,
"a {payload} byte payload under a {limit} byte limit has no room to expand"
);
}
}
#[test]
fn a_frame_of_wide_records_is_split_on_bytes() {
use surrealdb_types::prost::Message;
let width = *GRPC_MAX_MESSAGE_SIZE / QUERY_BATCH_RECORDS;
let records = 16 + 32 + 64 + 128 + QUERY_BATCH_RECORDS;
let rows =
Array::from((0..records).map(|_| Value::String("x".repeat(width))).collect::<Vec<_>>());
let frames = frames(vec![result(Value::Array(rows))]);
let batches: Vec<&rpc::QueryBatchFrame> =
frames[1..frames.len() - 1].iter().map(batch).collect();
let framed: usize = batches
.iter()
.map(|b| match &b.payload {
Some(rpc::query_batch_frame::Payload::Values(values)) => values.values.len(),
_ => 0,
})
.sum();
assert_eq!(framed, records, "every record still arrives, just in more frames");
assert!(
batches.len() > records.div_ceil(QUERY_BATCH_RECORDS),
"the record count alone would not have split this: {} frames",
batches.len()
);
let largest = frames.iter().map(|frame| frame.encoded_len()).max().expect("frames");
assert!(
largest + width >= frame_payload_budget(),
"the fixture has to fill a frame for this to test the budget: largest is {largest} \
against a {} byte budget",
frame_payload_budget()
);
for frame in &frames {
let encoded = frame.encoded_len();
assert!(
encoded + encoded / 256 <= *GRPC_MAX_MESSAGE_SIZE,
"a frame encodes to {encoded} bytes, which worst-case expansion takes past the \
{} byte limit both peers enforce",
*GRPC_MAX_MESSAGE_SIZE
);
}
}
#[test]
fn a_record_no_frame_can_carry_fails_its_statement() {
let rows = Array::from(vec![
Value::String("x".to_string()),
Value::String("x".repeat(*GRPC_MAX_MESSAGE_SIZE)),
]);
let frames = frames(vec![result(Value::Array(rows))]);
let batches: Vec<&rpc::QueryBatchFrame> =
frames[1..frames.len() - 1].iter().map(batch).collect();
let terminal = batches.last().expect("a terminal batch");
let error = terminal.error.as_ref().expect("the statement fails");
assert_eq!(error.kind, proto::ErrorKind::Validation as i32);
assert!(
error.message.contains("above the"),
"the refusal has to name the limit: {}",
error.message
);
assert!(
batches.iter().all(|b| b.payload.is_none()
|| matches!(
&b.payload,
Some(rpc::query_batch_frame::Payload::Values(v)) if v.values.is_empty()
)),
"a failed statement retracts its rows rather than framing them beside the refusal"
);
}
#[test]
fn an_oversized_error_is_trimmed_rather_than_lost() {
use surrealdb_types::prost::Message;
const BUDGET: usize = 512;
let thrown = TypesError::thrown("y".repeat(BUDGET * 4))
.with_cause(TypesError::internal("z".repeat(BUDGET * 4)));
let bounded = bounded_error(proto_error(&thrown), BUDGET);
assert!(
bounded.encoded_len() <= BUDGET,
"trimmed to {} bytes against a {BUDGET} byte budget",
bounded.encoded_len()
);
assert_eq!(
bounded.kind,
proto::ErrorKind::Thrown as i32,
"the kind is what a client branches on and survives every trim"
);
assert!(
bounded.message.ends_with(TRUNCATION_MARKER),
"a trimmed message says it was trimmed: {}",
bounded.message
);
assert!(bounded.cause.is_none(), "the cause chain goes before the message");
let small = proto_error(&TypesError::thrown("nope".to_string()));
assert_eq!(bounded_error(small.clone(), BUDGET), small);
}
#[test]
fn a_large_result_is_split_across_batches() {
let records = QUERY_BATCH_RECORDS * 2 + 1;
let rows = Array::from(
(0..records)
.map(|i| Value::Number(surrealdb_types::Number::Int(i as i64)))
.collect::<Vec<_>>(),
);
let frames = frames(vec![result(Value::Array(rows))]);
let batches: Vec<&rpc::QueryBatchFrame> =
frames[1..frames.len() - 1].iter().map(batch).collect();
assert!(batches.len() > 1, "a result this size must span several batches");
assert_eq!(
batches.iter().map(|b| b.batch_index).collect::<Vec<_>>(),
(0..batches.len() as u64).collect::<Vec<_>>(),
"batches must be indexed in order so the client can demultiplex them"
);
let kinds: Vec<i32> = batches.iter().map(|b| b.kind).collect();
let (last, rest) = kinds.split_last().expect("at least one batch");
assert_eq!(*last, rpc::QueryResponseKind::BatchedFinal as i32);
assert!(
rest.iter().all(|k| *k == rpc::QueryResponseKind::Batched as i32),
"only the last batch completes the statement"
);
let sizes = batch_sizes(&frames);
assert_eq!(sizes.iter().sum::<usize>(), records, "every record must be sent once");
assert_eq!(sizes[0], QUERY_FIRST_BATCH_RECORDS, "the first batch is small, for latency");
assert!(
sizes.windows(2).all(|w| w[1] >= w[0] || w[1] == sizes[sizes.len() - 1]),
"batches grow toward the cap, {sizes:?}"
);
assert!(sizes.iter().all(|n| *n <= QUERY_BATCH_RECORDS), "no batch exceeds the cap");
assert!(rest.iter().enumerate().all(|(i, _)| batches[i].stats.is_none()));
assert_eq!(
batches
.last()
.expect("at least one batch")
.stats
.as_ref()
.expect("the final batch carries the stats")
.records_returned,
records as i64
);
}
#[test]
fn an_empty_result_still_sends_one_final_batch() {
let frames = frames(vec![result(Value::Array(Array::new()))]);
assert_eq!(frames.len(), 3);
let batch = batch(&frames[1]);
assert_eq!(batch.kind, rpc::QueryResponseKind::BatchedFinal as i32);
assert_eq!(batch.stats.as_ref().expect("stats").records_returned, 0);
}
#[test]
fn list_and_scalar_results_use_distinct_kinds() {
let list = frames(vec![result(Value::Array(Array::from(vec![
Value::Bool(true),
Value::Bool(false),
])))]);
let list = batch(&list[1]);
assert_eq!(list.kind, rpc::QueryResponseKind::BatchedFinal as i32);
assert_eq!(list.stats.as_ref().expect("stats").records_returned, 2);
let scalar = frames(vec![result(Value::Bool(true))]);
let scalar = batch(&scalar[1]);
assert_eq!(scalar.kind, rpc::QueryResponseKind::Single as i32);
assert_eq!(scalar.stats.as_ref().expect("stats").records_returned, 1);
}
#[test]
fn a_failed_statement_does_not_end_the_stream() {
let frames = frames(vec![
QueryResult {
time: Duration::ZERO,
result: Err(TypesError::query("boom".to_string(), None)),
query_type: QueryType::Other,
},
result(Value::Bool(true)),
]);
let failed = batch(&frames[1]);
assert_eq!(failed.error.as_ref().expect("error").kind, proto::ErrorKind::Query as i32);
assert!(failed.payload.is_none());
assert!(batch(&frames[2]).error.is_none());
assert!(matches!(frames[3].frame, Some(rpc::query_response::Frame::End(_))));
}
#[test]
fn live_and_kill_statements_are_labelled() {
let live = frames(vec![QueryResult {
time: Duration::ZERO,
result: Ok(Value::None),
query_type: QueryType::Live,
}]);
assert_eq!(batch(&live[1]).statement_kind, rpc::QueryStatementKind::Live as i32);
let kill = frames(vec![QueryResult {
time: Duration::ZERO,
result: Ok(Value::None),
query_type: QueryType::Kill,
}]);
assert_eq!(batch(&kill[1]).statement_kind, rpc::QueryStatementKind::Kill as i32);
}
#[test]
fn nullable_strings_are_three_state() {
assert_eq!(from_nullable(None), Value::None);
assert_eq!(
from_nullable(Some(rpc::NullableString {
value: Some(rpc::nullable_string::Value::Null(proto::NullValue {})),
})),
Value::Null
);
assert_eq!(
from_nullable(Some(rpc::NullableString {
value: Some(rpc::nullable_string::Value::Some("test".to_string())),
})),
Value::String("test".to_string())
);
}
#[test]
fn user_credentials_flatten_to_the_signin_object() {
let object = access_credentials(rpc::AccessMethod {
method: Some(rpc::access_method::Method::User(rpc::UserCredentials {
namespace: String::new(),
database: String::new(),
username: "root".to_string(),
password: "root".to_string(),
access: String::new(),
})),
})
.expect("credentials");
assert_eq!(object.get("user"), Some(&Value::String("root".to_string())));
assert_eq!(object.get("pass"), Some(&Value::String("root".to_string())));
assert!(object.get("ns").is_none());
assert!(object.get("db").is_none());
}
#[test]
fn record_credentials_carry_their_variables() {
let variables: proto::Object = std::collections::BTreeMap::from([(
"email".to_string(),
proto::Value::try_from(Value::String("a@b.c".to_string())).expect("encodable"),
)])
.into();
let object = record_credentials(rpc::RecordCredentials {
namespace: "test".to_string(),
database: "test".to_string(),
access: "user".to_string(),
variables: Some(proto::Variables {
variables: variables.items,
}),
})
.expect("credentials");
assert_eq!(object.get("ns"), Some(&Value::String("test".to_string())));
assert_eq!(object.get("ac"), Some(&Value::String("user".to_string())));
assert_eq!(object.get("email"), Some(&Value::String("a@b.c".to_string())));
}
#[test]
fn export_table_selection_maps_each_arm() {
let config = |tables: Option<rpc::export_config::Tables>| {
from_proto_export_config(rpc::ExportConfig {
tables,
..Default::default()
})
.tables
};
assert!(matches!(config(None), export::TableConfig::All));
assert!(matches!(
config(Some(rpc::export_config::Tables::from(false))),
export::TableConfig::None
));
match config(Some(rpc::export_config::Tables {
selection: Some(rpc::export_config::tables::Selection::Selected(
rpc::export_config::SelectedTables {
tables: vec!["person".to_string()],
},
)),
})) {
export::TableConfig::Some(tables) => assert_eq!(tables, vec!["person".to_string()]),
other => panic!("expected a table selection, got {other:?}"),
}
}
#[test]
fn error_kinds_map_to_round_tripping_status_codes() {
use tonic::Code;
let cases = [
(TypesError::validation("v".to_string(), None), Code::InvalidArgument),
(TypesError::configuration("c".to_string(), None), Code::Unimplemented),
(TypesError::query("q".to_string(), None), Code::Aborted),
(TypesError::not_allowed("n".to_string(), None), Code::PermissionDenied),
(TypesError::not_found("f".to_string(), None), Code::NotFound),
(TypesError::already_exists("a".to_string(), None), Code::AlreadyExists),
(TypesError::connection("x".to_string(), None), Code::Unavailable),
(TypesError::internal("i".to_string()), Code::Internal),
];
for (error, code) in cases {
assert_eq!(to_status(&error).code(), code, "{}", error.message());
}
}
#[test]
fn proto_errors_keep_their_cause_chain() {
let error = TypesError::query("outer".to_string(), None).with_cause(
TypesError::serialization("inner".to_string(), SerializationError::Deserialization),
);
let proto = to_proto_error(&error);
assert_eq!(proto.kind, proto::ErrorKind::Query as i32);
let cause = proto.cause.expect("cause");
assert_eq!(cause.kind, proto::ErrorKind::Serialization as i32);
assert_eq!(cause.message, "inner");
}
#[test]
fn proto_errors_carry_their_reason() {
let nested = to_proto_error(&TypesError::not_allowed(
"nope".to_string(),
NotAllowedError::Auth(AuthError::TokenExpired),
));
assert_eq!(nested.kind, proto::ErrorKind::NotAllowed as i32);
assert_eq!(nested.detail_kind(), Some("Auth"));
let content = nested.details.and_then(|details| details.content).expect("nested content");
assert_eq!(
Value::try_from(content).expect("a decodable payload"),
AuthError::TokenExpired.into_value()
);
let timed_out = to_proto_error(&TypesError::query(
"too slow".to_string(),
QueryError::TimedOut {
duration: Duration::from_secs(5),
},
));
assert_eq!(timed_out.detail_kind(), Some("TimedOut"));
assert!(timed_out.details.and_then(|details| details.content).is_some());
assert_eq!(to_proto_error(&TypesError::internal("i".to_string())).detail_kind(), None);
}
#[test]
fn a_conflict_is_the_only_error_marked_retryable() {
let conflict = to_proto_error(&TypesError::query(
"This transaction can be retried".to_string(),
QueryError::TransactionConflict,
));
assert!(conflict.is_retryable());
assert_eq!(conflict.detail_kind(), Some("TransactionConflict"));
for error in [
TypesError::query("cancelled".to_string(), QueryError::Cancelled),
TypesError::query("no reason given".to_string(), None),
TypesError::internal("broken".to_string()),
] {
assert!(!to_proto_error(&error).is_retryable(), "{}", error.message());
}
}
#[tokio::test]
async fn the_advertised_message_size_is_the_configured_one() {
let limits = service().await.server_capabilities().limits.expect("limits");
assert_eq!(limits.max_message_bytes, *GRPC_MAX_MESSAGE_SIZE as u64);
}
#[tokio::test]
async fn the_smallest_allowed_message_size_still_carries_the_handshake() {
use surrealdb_types::prost::Message;
let capabilities = service().await.server_capabilities();
assert!(
capabilities.encoded_len() < crate::cnf::GRPC_MIN_MESSAGE_SIZE,
"the handshake encodes to {} bytes, which the {} byte floor cannot carry",
capabilities.encoded_len(),
crate::cnf::GRPC_MIN_MESSAGE_SIZE
);
}
#[test]
fn tokens_are_read_from_either_answer_shape() {
let bare = to_tokens(DbResult::Other(Value::String("access".to_string()))).expect("tokens");
assert_eq!(bare.access, "access");
assert!(bare.refresh.is_empty());
let mut object = Object::new();
object.insert("access".to_string(), Value::String("access".to_string()));
object.insert("refresh".to_string(), Value::String("refresh".to_string()));
let pair = to_tokens(DbResult::Other(Value::Object(object))).expect("tokens");
assert_eq!(pair.access, "access");
assert_eq!(pair.refresh, "refresh");
}
#[test]
fn trailers_are_verified_against_what_arrived() {
let digest = blake3::hash(b"hello");
let trailer = rpc::DataTrailer {
bytes: 5,
blake3: digest.to_hex().to_string(),
};
assert!(verify_trailer(&trailer, 5, digest).is_ok());
assert!(verify_trailer(&trailer, 4, digest).is_err());
let unchecked = rpc::DataTrailer {
bytes: 5,
blake3: String::new(),
};
assert!(verify_trailer(&unchecked, 5, digest).is_ok());
let wrong = rpc::DataTrailer {
bytes: 5,
blake3: blake3::hash(b"other").to_hex().to_string(),
};
assert!(verify_trailer(&wrong, 5, digest).is_err());
}
}