use std::path::PathBuf;
use std::sync::{Arc, Mutex};
use std::time::Duration as StdDuration;
use async_channel::Sender;
use surrealdb_engine_api::{
DbExportConfig, EngineContext, EngineFuture, MlExportConfig, SurrealEngine,
};
use surrealdb_protocol::proto::rpc::v1 as rpc;
use surrealdb_protocol::proto::rpc::v1::surreal_db_service_client::SurrealDbServiceClient;
use surrealdb_protocol::proto::v1 as proto;
use surrealdb_rpc::{QueryResult, QueryStreamItem, QueryType, Token};
use tonic::codec::CompressionEncoding;
use tonic::transport::Channel;
use uuid::Uuid;
use crate::types::prost::Message;
use crate::types::{
Action, Array, Notification, Object, SerializationError, ValidationError, Value, Variables,
};
use crate::{Error, ExtraFeatures, SessionId};
type EngineResult<T> = crate::Result<T>;
#[derive(Debug)]
pub struct Grpc;
#[derive(Debug)]
pub struct Grpcs;
#[derive(Debug, Clone)]
pub struct Client(());
impl crate::Connection for Client {}
impl crate::conn::Sealed for Client {
#[allow(private_interfaces)]
fn connect(
address: crate::opt::Endpoint,
_capacity: usize,
session_clone: Option<crate::SessionClone>,
) -> crate::method::BoxFuture<'static, crate::Result<crate::Surreal<Self>>> {
Box::pin(async move {
let session_clone = session_clone.unwrap_or_else(crate::SessionClone::new);
let (engine, features) =
connect_engine(&address, session_clone.receiver.clone()).await?;
let router = crate::conn::Router::from_engine(engine, features, address.config);
let waiter = tokio::sync::watch::channel(Some(crate::opt::WaitFor::Connection));
Ok((router, waiter, session_clone).into())
})
}
}
pub(crate) async fn connect_engine(
address: &crate::opt::Endpoint,
session_rx: async_channel::Receiver<SessionId>,
) -> crate::Result<(Arc<dyn SurrealEngine>, std::collections::HashSet<ExtraFeatures>)> {
let is_tls = address.url.scheme() == "grpcs";
let scheme = if is_tls {
"https"
} else {
"http"
};
let dst = match address.url.as_str().split_once("://") {
Some((_, rest)) => format!("{scheme}://{rest}"),
None => {
return Err(Error::configuration(
format!("Invalid gRPC endpoint: {}", address.url),
None,
));
}
};
#[cfg_attr(not(feature = "rustls"), expect(unused_mut))]
let mut builder = tonic::transport::Endpoint::from_shared(dst)
.map_err(|e| Error::configuration(e.to_string(), None))?;
#[cfg(any(feature = "native-tls", feature = "rustls"))]
if address.config.tls_config.is_some() {
return Err(Error::configuration(
"A custom TLS configuration is not supported over `grpcs://`; \
the connection would fall back to the public roots instead"
.to_string(),
None,
));
}
if is_tls {
#[cfg(feature = "rustls")]
{
builder = builder
.tls_config(tonic::transport::ClientTlsConfig::new().with_webpki_roots())
.map_err(|e| Error::configuration(e.to_string(), None))?;
}
#[cfg(not(feature = "rustls"))]
return Err(Error::configuration(
"Connecting over `grpcs://` requires the `rustls` feature".to_string(),
None,
));
}
let channel = builder.connect().await.map_err(|e| {
Error::connection(e.to_string(), crate::types::ConnectionError::ConnectionFailed)
})?;
let mut client = SurrealDbServiceClient::new(channel)
.accept_compressed(CompressionEncoding::Zstd)
.accept_compressed(CompressionEncoding::Gzip);
if let Some(limit) = address.config.grpc.max_message_size {
client = client.max_decoding_message_size(limit).max_encoding_message_size(limit);
}
let capabilities = fetch_capabilities(&mut client).await?;
let compress_requests = server_accepts_request_encoding(&capabilities);
let advertised = capabilities
.limits
.as_ref()
.map(|limits| limits.max_message_bytes)
.and_then(|limit| usize::try_from(limit).ok())
.filter(|limit| *limit > 0);
let configured = address.config.grpc.max_message_size;
if let (Some(configured), Some(advertised)) = (configured, advertised)
&& configured > advertised
{
warn!(
"The configured gRPC message size ({configured} bytes) is larger than the server \
accepts ({advertised} bytes); requests stay bounded by the server's figure. Raise \
SURREAL_GRPC_MAX_MESSAGE_SIZE on the server to send more than that."
);
}
if let Some(limit) = max_response_size(configured, advertised) {
client = client.max_decoding_message_size(limit);
}
let max_request_size = max_request_size(configured, advertised);
if let Some((limit, _)) = max_request_size {
client = client.max_encoding_message_size(limit);
}
let features = extra_features(&capabilities);
let engine = Arc::new(GrpcEngine {
client,
compress_requests,
max_request_size,
server_version: capabilities.server_version,
sessions: SessionRegistry::default(),
query_timeout: address
.config
.query_timeout
.map(proto::Duration::try_from)
.transpose()
.map_err(|e| Error::configuration(format!("Invalid query timeout: {e}"), None))?,
});
tokio::spawn(session_task(Arc::clone(&engine), session_rx));
Ok((engine, features))
}
fn extra_features(
capabilities: &rpc::ServerCapabilities,
) -> std::collections::HashSet<ExtraFeatures> {
const EXPORT: &str = "surrealdb.protocol.rpc.v1.SurrealDBService/ExportSurql";
let denied = |method: &str| capabilities.denied_methods.iter().any(|m| m == method);
let mut features = std::collections::HashSet::new();
if !denied(EXPORT) {
features.insert(ExtraFeatures::Backup);
}
let reports_any = !capabilities.capabilities.is_empty();
if !reports_any || capabilities.capabilities.iter().any(|c| c == "LIVE_QUERIES") {
features.insert(ExtraFeatures::LiveQueries);
}
features
}
const REQUEST_ENCODING: CompressionEncoding = CompressionEncoding::Zstd;
const REQUEST_ENCODING_TOKEN: &str = "zstd";
fn server_accepts_request_encoding(capabilities: &rpc::ServerCapabilities) -> bool {
capabilities
.accepted_message_encodings
.iter()
.any(|encoding| encoding.eq_ignore_ascii_case(REQUEST_ENCODING_TOKEN))
}
fn rejected_request_encoding(status: &tonic::Status) -> bool {
status.code() == tonic::Code::Unimplemented
&& status.metadata().contains_key("grpc-accept-encoding")
}
async fn fetch_capabilities(
client: &mut SurrealDbServiceClient<Channel>,
) -> crate::Result<rpc::ServerCapabilities> {
get_capabilities(client).await.map_err(|status| {
Error::connection(
status.message().to_string(),
crate::types::ConnectionError::ConnectionFailed,
)
})
}
async fn get_capabilities(
client: &mut SurrealDbServiceClient<Channel>,
) -> Result<rpc::ServerCapabilities, tonic::Status> {
let request = rpc::GetCapabilitiesRequest {
context: None,
client: Some(rpc::ClientInfo {
name: "surrealdb-rust".to_string(),
version: env!("CARGO_PKG_VERSION").to_string(),
platform: "rust".to_string(),
metadata: Vec::new(),
}),
};
client
.get_capabilities(request)
.await?
.into_inner()
.capabilities
.ok_or_else(|| tonic::Status::internal("Server did not report its capabilities"))
}
#[derive(Debug, Clone)]
enum Replayable {
Use {
namespace: Option<String>,
database: Option<String>,
},
Set {
key: String,
value: Value,
},
Unset {
key: String,
},
Signin(Object),
Signup(Object),
Authenticate(Token),
Invalidate,
}
impl Replayable {
fn supersedes(&self, prev: &Self) -> bool {
match (prev, self) {
(
Replayable::Use {
namespace: pn,
database: pd,
},
Replayable::Use {
namespace: nn,
database: nd,
},
) => (pn.is_none() || nn.is_some()) && (pd.is_none() || nd.is_some()),
(
Replayable::Set {
key: previous,
..
}
| Replayable::Unset {
key: previous,
},
Replayable::Set {
key: next,
..
}
| Replayable::Unset {
key: next,
},
) => previous == next,
_ => false,
}
}
fn resets_session_state(&self) -> bool {
matches!(
self,
Replayable::Signin(_)
| Replayable::Signup(_)
| Replayable::Authenticate(_)
| Replayable::Invalidate
)
}
}
#[derive(Debug)]
struct GrpcSession {
server: Uuid,
replay: Mutex<Vec<Replayable>>,
}
impl GrpcSession {
fn new(server: Uuid) -> Arc<Self> {
Arc::new(Self {
server,
replay: Mutex::new(Vec::new()),
})
}
fn record(&self, op: Replayable) {
let mut replay = self.replay.lock().expect("session registry poisoned");
for previous in replay.iter_mut().rev() {
if previous.resets_session_state() {
break;
}
if op.supersedes(previous) {
*previous = op;
return;
}
}
replay.push(op);
}
fn log(&self) -> Vec<Replayable> {
self.replay.lock().expect("session registry poisoned").clone()
}
}
type SessionRegistry = surrealdb_engine_api::SessionRegistry<Arc<GrpcSession>, Error>;
struct GrpcEngine {
client: SurrealDbServiceClient<Channel>,
compress_requests: bool,
server_version: String,
sessions: SessionRegistry,
max_request_size: Option<(usize, LimitSource)>,
query_timeout: Option<proto::Duration>,
}
impl std::fmt::Debug for GrpcEngine {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("GrpcEngine").field("server_version", &self.server_version).finish()
}
}
impl GrpcEngine {
fn client(&self) -> SurrealDbServiceClient<Channel> {
self.client.clone()
}
fn checked<M: Message>(&self, request: M) -> EngineResult<M> {
check_message_size(self.max_request_size, &request)?;
Ok(request)
}
fn compressed_upload_client(&self) -> Option<SurrealDbServiceClient<Channel>> {
self.compress_requests.then(|| self.client.clone().send_compressed(REQUEST_ENCODING))
}
fn context(&self, session: Uuid, transaction: Option<Uuid>) -> rpc::RequestContext {
rpc::RequestContext {
session: Some(proto::Uuid::from_uuid(session)),
transaction: transaction.map(proto::Uuid::from_uuid),
timeout: self.query_timeout,
}
}
async fn ready(&self, ctx: EngineContext) -> EngineResult<rpc::RequestContext> {
Ok(self.ready_session(ctx).await?.0)
}
async fn ready_session(
&self,
ctx: EngineContext,
) -> EngineResult<(rpc::RequestContext, Arc<GrpcSession>)> {
let session = self.sessions.resolve(ctx.session).await?;
let context = self.context(session.server, ctx.transaction);
Ok((context, session))
}
async fn establish_clone(&self, log: Vec<Replayable>) -> EngineResult<Arc<GrpcSession>> {
let session = GrpcSession::new(self.attach().await?);
for op in log {
self.apply(self.context(session.server, None), op.clone()).await?;
session.record(op);
}
Ok(session)
}
fn publish(&self, id: Uuid, established: EngineResult<Arc<GrpcSession>>) {
if let Err(error) = established.as_ref() {
trace!("failed to establish session {id}: {error}");
}
self.sessions.entry(id).publish(established);
}
async fn apply(&self, context: rpc::RequestContext, op: Replayable) -> EngineResult<()> {
match op {
Replayable::Use {
namespace,
database,
} => {
self.use_ns_db_inner(context, namespace, database).await?;
}
Replayable::Set {
key,
value,
} => self.set_inner(context, key, value).await?,
Replayable::Unset {
key,
} => self.unset_inner(context, key).await?,
Replayable::Signin(credentials) => {
self.signin_inner(context, credentials).await?;
}
Replayable::Signup(credentials) => {
self.signup_inner(context, credentials).await?;
}
Replayable::Authenticate(token) => {
self.authenticate_inner(context, token).await?;
}
Replayable::Invalidate => self.invalidate_inner(context).await?,
}
Ok(())
}
async fn attach(&self) -> EngineResult<Uuid> {
let response = self
.client()
.attach_session(self.checked(rpc::AttachSessionRequest {
context: None,
})?)
.await
.map_err(status_to_error)?
.into_inner();
response
.session
.ok_or_else(|| Error::internal("Server did not allocate a session".to_string()))?
.to_uuid()
.map_err(|e| Error::internal(e.to_string()))
}
async fn detach(&self, session: Uuid) -> EngineResult<()> {
self.client()
.detach_session(self.checked(rpc::DetachSessionRequest {
context: Some(self.context(session, None)),
})?)
.await
.map_err(status_to_error)?;
Ok(())
}
async fn use_ns_db_inner(
&self,
context: rpc::RequestContext,
namespace: Option<String>,
database: Option<String>,
) -> EngineResult<(Option<String>, Option<String>)> {
let response = self
.client()
.r#use(self.checked(rpc::UseRequest {
context: Some(context),
namespace: namespace.map(nullable),
database: database.map(nullable),
})?)
.await
.map_err(status_to_error)?
.into_inner();
let opt = |s: String| (!s.is_empty()).then_some(s);
Ok((opt(response.namespace), opt(response.database)))
}
async fn set_inner(
&self,
context: rpc::RequestContext,
key: String,
value: Value,
) -> EngineResult<()> {
let request = rpc::SetVariableRequest {
context: Some(context),
name: key,
value: Some(to_proto_value(value)?),
};
self.client().set_variable(self.checked(request)?).await.map_err(status_to_error)?;
Ok(())
}
async fn unset_inner(&self, context: rpc::RequestContext, key: String) -> EngineResult<()> {
self.client()
.unset_variable(self.checked(rpc::UnsetVariableRequest {
context: Some(context),
name: key,
})?)
.await
.map_err(status_to_error)?;
Ok(())
}
async fn signin_inner(
&self,
context: rpc::RequestContext,
credentials: Object,
) -> EngineResult<Token> {
let request = rpc::SigninRequest {
context: Some(context),
access_method: Some(access_method(credentials)?),
};
let response = self
.client()
.signin(self.checked(request)?)
.await
.map_err(status_to_error)?
.into_inner();
tokens_to_token(response.tokens)
}
async fn signup_inner(
&self,
context: rpc::RequestContext,
credentials: Object,
) -> EngineResult<Token> {
let request = rpc::SignupRequest {
context: Some(context),
credentials: Some(record_credentials(credentials)?),
};
let response = self
.client()
.signup(self.checked(request)?)
.await
.map_err(status_to_error)?
.into_inner();
tokens_to_token(response.tokens)
}
async fn authenticate_inner(
&self,
context: rpc::RequestContext,
token: Token,
) -> EngineResult<Token> {
let access = match &token {
Token::Access(access) => access.clone(),
Token::WithRefresh {
access,
..
} => access.clone(),
};
let response = self
.client()
.authenticate(self.checked(rpc::AuthenticateRequest {
context: Some(context),
token: access,
})?)
.await
.map_err(status_to_error)?
.into_inner();
match response.tokens {
Some(_) => tokens_to_token(response.tokens),
None => Ok(token),
}
}
async fn invalidate_inner(&self, context: rpc::RequestContext) -> EngineResult<()> {
self.client()
.invalidate(self.checked(rpc::InvalidateRequest {
context: Some(context),
})?)
.await
.map_err(status_to_error)?;
Ok(())
}
async fn open_query(
&self,
context: rpc::RequestContext,
query: String,
variables: Variables,
max_batch_records: u32,
) -> EngineResult<tonic::Streaming<rpc::QueryResponse>> {
let request = rpc::QueryRequest {
context: Some(context),
query,
variables: Some(to_proto_variables(variables)?),
accepted_encodings: Vec::new(),
max_batch_records,
};
Ok(self.client().query(self.checked(request)?).await.map_err(status_to_error)?.into_inner())
}
async fn query_inner(
&self,
context: rpc::RequestContext,
query: String,
variables: Variables,
) -> EngineResult<Vec<QueryResult>> {
let mut stream = self.open_query(context, query, variables, 0).await?;
let mut statements: Vec<Option<Statement>> = Vec::new();
let mut ended = None;
while let Some(response) = stream.message().await.map_err(status_to_error)? {
match response.frame {
Some(rpc::query_response::Frame::Begin(begin)) => {
statements = Vec::new();
grow_statements(&mut statements, begin.statement_count as usize)?;
}
Some(rpc::query_response::Frame::Batch(batch)) => {
let index = batch.query_index as usize;
if index >= statements.len() {
grow_statements(&mut statements, index + 1)?;
}
statements[index].get_or_insert_with(Statement::default).push(batch);
}
Some(rpc::query_response::Frame::End(end)) => {
ended = Some(end);
break;
}
Some(rpc::query_response::Frame::Error(error)) => {
return Err(proto_error(error));
}
None => {
return Err(Error::internal(
"Query stream carried an unrecognised frame".to_string(),
));
}
}
}
let Some(end) = ended else {
return Err(Error::connection(
"The query ended before it was complete".to_string(),
crate::types::ConnectionError::ConnectionFailed,
));
};
let results = finish_statements(statements);
if results.len() != end.result_count as usize {
return Err(Error::connection(
format!(
"The query reported {} statement results but {} arrived",
end.result_count,
results.len()
),
crate::types::ConnectionError::ConnectionFailed,
));
}
Ok(results)
}
}
fn deserialization_error(error: impl std::fmt::Display) -> Error {
Error::serialization(error.to_string(), SerializationError::Deserialization)
}
const MAX_STATEMENTS: usize = 1 << 20;
fn grow_statements(statements: &mut Vec<Option<Statement>>, len: usize) -> EngineResult<()> {
if len > MAX_STATEMENTS {
return Err(Error::internal(format!(
"The server reported {len} statement results, more than the {MAX_STATEMENTS} a query may have"
)));
}
statements.resize_with(len, || None);
Ok(())
}
async fn stream_query(
engine: &GrpcEngine,
context: rpc::RequestContext,
query: String,
variables: Variables,
max_batch_records: u32,
items: Sender<QueryStreamItem>,
) -> EngineResult<()> {
let mut stream = engine.open_query(context, query, variables, max_batch_records).await?;
let mut ended = false;
while let Some(response) = stream.message().await.map_err(status_to_error)? {
let batch = match response.frame {
Some(rpc::query_response::Frame::Batch(batch)) => batch,
Some(rpc::query_response::Frame::End(_)) => {
ended = true;
break;
}
Some(rpc::query_response::Frame::Error(error)) => return Err(proto_error(error)),
Some(rpc::query_response::Frame::Begin(_)) => continue,
None => {
return Err(Error::internal(
"Query stream carried an unrecognised frame".to_string(),
));
}
};
for item in batch_items(batch)? {
if items.send(item).await.is_err() {
return Ok(());
}
}
}
if !ended {
return Err(Error::connection(
"The query ended before it was complete".to_string(),
crate::types::ConnectionError::ConnectionFailed,
));
}
Ok(())
}
fn batch_items(batch: rpc::QueryBatchFrame) -> EngineResult<Vec<QueryStreamItem>> {
let index = batch.query_index as usize;
let single = batch.kind == rpc::QueryResponseKind::Single as i32;
let terminal = single || batch.kind == rpc::QueryResponseKind::BatchedFinal as i32;
let mut items = Vec::new();
if let Some(error) = batch.error {
items.push(QueryStreamItem::Finished {
index,
time: batch_duration(batch.stats.as_ref()),
query_type: statement_query_type(batch.statement_kind),
error: Some(proto_error(error)),
});
return Ok(items);
}
let values = match batch.payload {
Some(rpc::query_batch_frame::Payload::Values(values)) => values
.values
.into_iter()
.map(|value| Value::try_from(value).map_err(deserialization_error))
.collect::<EngineResult<Vec<Value>>>()?,
Some(rpc::query_batch_frame::Payload::Arrow(_)) => {
return Err(Error::internal(
"Server sent a columnar batch, which was not requested".to_string(),
));
}
None => Vec::new(),
};
if single {
if let Some(value) = values.into_iter().next() {
items.push(QueryStreamItem::Value {
index,
value,
});
}
} else if !values.is_empty() {
items.push(QueryStreamItem::Rows {
index,
values,
});
}
if terminal {
items.push(QueryStreamItem::Finished {
index,
time: batch_duration(batch.stats.as_ref()),
query_type: statement_query_type(batch.statement_kind),
error: None,
});
}
Ok(items)
}
fn batch_duration(stats: Option<&rpc::QueryStats>) -> StdDuration {
stats
.and_then(|stats| stats.execution_duration)
.and_then(|duration| StdDuration::try_from(duration).ok())
.unwrap_or_default()
}
fn statement_query_type(kind: i32) -> QueryType {
match rpc::QueryStatementKind::try_from(kind) {
Ok(rpc::QueryStatementKind::Live) => QueryType::Live,
Ok(rpc::QueryStatementKind::Kill) => QueryType::Kill,
_ => QueryType::Other,
}
}
fn finish_statements(statements: Vec<Option<Statement>>) -> Vec<QueryResult> {
statements.into_iter().flatten().map(Statement::finish).collect()
}
async fn session_task(engine: Arc<GrpcEngine>, session_rx: async_channel::Receiver<SessionId>) {
while let Ok(event) = session_rx.recv().await {
match event {
SessionId::Initial(id) => {
let established = engine.attach().await.map(GrpcSession::new);
engine.publish(id, established);
}
SessionId::Clone {
old,
new,
} => {
let log = match engine.sessions.established(old) {
Some(Ok(session)) => session.log(),
_ => Vec::new(),
};
engine.publish(new, engine.establish_clone(log).await);
}
SessionId::Drop(id) => {
match engine.sessions.end(id) {
Some(Ok(session)) => {
if let Err(error) = engine.detach(session.server).await {
trace!("failed to detach session {id}: {error}");
}
}
Some(Err(error)) => trace!("session {id} was never attached: {error}"),
None => trace!("session {id} was dropped before it was established"),
}
}
}
}
}
impl SurrealEngine for GrpcEngine {
fn query(
&self,
ctx: EngineContext,
query: std::borrow::Cow<'static, str>,
variables: Variables,
) -> EngineFuture<'_, Vec<QueryResult>> {
Box::pin(async move {
let context = self.ready(ctx).await?;
self.query_inner(context, query.into_owned(), variables).await
})
}
fn query_stream(
&self,
ctx: EngineContext,
query: std::borrow::Cow<'static, str>,
variables: Variables,
items: Sender<QueryStreamItem>,
) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
stream_query(self, context, query.into_owned(), variables, 0, items).await
})
}
fn run(
&self,
ctx: EngineContext,
name: String,
version: Option<String>,
args: Array,
) -> EngineFuture<'_, Value> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let args = args.into_iter().map(to_proto_value).collect::<EngineResult<Vec<_>>>()?;
let request = rpc::RunRequest {
context: Some(context),
name,
version: version.unwrap_or_default(),
args,
};
let response = self
.client()
.run(self.checked(request)?)
.await
.map_err(status_to_error)?
.into_inner();
match response.result {
Some(value) => Value::try_from(value).map_err(deserialization_error),
None => Ok(Value::None),
}
})
}
fn use_ns_db(
&self,
ctx: EngineContext,
namespace: Option<String>,
database: Option<String>,
) -> EngineFuture<'_, (Option<String>, Option<String>)> {
Box::pin(async move {
let (context, session) = self.ready_session(ctx).await?;
let selection =
self.use_ns_db_inner(context, namespace.clone(), database.clone()).await?;
session.record(Replayable::Use {
namespace,
database,
});
Ok(selection)
})
}
fn set(&self, ctx: EngineContext, key: String, value: Value) -> EngineFuture<'_, ()> {
Box::pin(async move {
let (context, session) = self.ready_session(ctx).await?;
self.set_inner(context, key.clone(), value.clone()).await?;
session.record(Replayable::Set {
key,
value,
});
Ok(())
})
}
fn unset(&self, ctx: EngineContext, key: String) -> EngineFuture<'_, ()> {
Box::pin(async move {
let (context, session) = self.ready_session(ctx).await?;
self.unset_inner(context, key.clone()).await?;
session.record(Replayable::Unset {
key,
});
Ok(())
})
}
fn signup(&self, ctx: EngineContext, credentials: Object) -> EngineFuture<'_, Token> {
Box::pin(async move {
let (context, session) = self.ready_session(ctx).await?;
let token = self.signup_inner(context, credentials.clone()).await?;
session.record(Replayable::Signup(credentials));
Ok(token)
})
}
fn signin(&self, ctx: EngineContext, credentials: Object) -> EngineFuture<'_, Token> {
Box::pin(async move {
let (context, session) = self.ready_session(ctx).await?;
let token = self.signin_inner(context, credentials.clone()).await?;
session.record(Replayable::Signin(credentials));
Ok(token)
})
}
fn authenticate(&self, ctx: EngineContext, token: Token) -> EngineFuture<'_, Token> {
Box::pin(async move {
let (context, session) = self.ready_session(ctx).await?;
let token = self.authenticate_inner(context, token).await?;
session.record(Replayable::Authenticate(token.clone()));
Ok(token)
})
}
fn refresh(&self, ctx: EngineContext, token: Token) -> EngineFuture<'_, Token> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let Token::WithRefresh {
access,
refresh,
} = token
else {
return Err(Error::validation(
"This token carries no refresh token".to_string(),
None,
));
};
let response = self
.client()
.refresh_tokens(self.checked(rpc::RefreshTokensRequest {
context: Some(context),
access,
refresh,
})?)
.await
.map_err(status_to_error)?
.into_inner();
tokens_to_token(response.tokens)
})
}
fn revoke(&self, ctx: EngineContext, token: Token) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let (access, refresh) = match token {
Token::Access(access) => (access, String::new()),
Token::WithRefresh {
access,
refresh,
} => (access, refresh),
};
self.client()
.revoke_tokens(self.checked(rpc::RevokeTokensRequest {
context: Some(context),
access,
refresh,
})?)
.await
.map_err(status_to_error)?;
Ok(())
})
}
fn invalidate(&self, ctx: EngineContext) -> EngineFuture<'_, ()> {
Box::pin(async move {
let (context, session) = self.ready_session(ctx).await?;
self.invalidate_inner(context).await?;
session.record(Replayable::Invalidate);
Ok(())
})
}
fn begin(&self, ctx: EngineContext) -> EngineFuture<'_, Uuid> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let response = self
.client()
.begin_transaction(self.checked(rpc::BeginTransactionRequest {
context: Some(context),
})?)
.await
.map_err(status_to_error)?
.into_inner();
response
.transaction
.ok_or_else(|| {
Error::internal("Server did not return a transaction id".to_string())
})?
.to_uuid()
.map_err(|e| Error::internal(e.to_string()))
})
}
fn commit(&self, ctx: EngineContext, txn: Uuid) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context =
self.ready(EngineContext::with_transaction(ctx.session, Some(txn))).await?;
self.client()
.commit_transaction(self.checked(rpc::CommitTransactionRequest {
context: Some(context),
})?)
.await
.map_err(status_to_error)?;
Ok(())
})
}
fn rollback(&self, ctx: EngineContext, txn: Uuid) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context =
self.ready(EngineContext::with_transaction(ctx.session, Some(txn))).await?;
self.client()
.cancel_transaction(self.checked(rpc::CancelTransactionRequest {
context: Some(context),
})?)
.await
.map_err(status_to_error)?;
Ok(())
})
}
fn health(&self, ctx: EngineContext) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
self.client()
.health(self.checked(rpc::HealthRequest {
context: Some(context),
})?)
.await
.map_err(status_to_error)?;
Ok(())
})
}
fn version(&self, _ctx: EngineContext) -> EngineFuture<'_, String> {
Box::pin(async move { Ok(self.server_version.clone()) })
}
fn kill(&self, ctx: EngineContext, uuid: Uuid) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
self.client()
.kill(self.checked(rpc::KillRequest {
context: Some(context),
live_query_id: Some(proto::Uuid::from_uuid(uuid)),
})?)
.await
.map_err(status_to_error)?;
Ok(())
})
}
fn subscribe_live(
&self,
ctx: EngineContext,
uuid: Uuid,
notifications: async_channel::Sender<Result<Notification, Error>>,
) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let stream = self
.client()
.subscribe(self.checked(rpc::SubscribeRequest {
context: Some(context),
resume_from: None,
subscribe_to: Some(rpc::subscribe_request::SubscribeTo::LiveQueryId(
proto::Uuid::from_uuid(uuid),
)),
})?)
.await
.map_err(status_to_error)?
.into_inner();
tokio::spawn(pump_notifications(stream, notifications, uuid));
Ok(())
})
}
fn export_file(
&self,
ctx: EngineContext,
path: PathBuf,
config: Option<DbExportConfig>,
) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let stream = self.export_surql(context, config).await?;
write_export_to_file(stream, path).await
})
}
fn export_bytes(
&self,
ctx: EngineContext,
bytes: async_channel::Sender<Result<Vec<u8>, Error>>,
config: Option<DbExportConfig>,
) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let stream = self.export_surql(context, config).await?;
tokio::spawn(pump_export_to_channel(stream, bytes));
Ok(())
})
}
fn export_ml_file(
&self,
ctx: EngineContext,
path: PathBuf,
config: MlExportConfig,
) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let stream = self.export_ml(context, config).await?;
write_export_to_file(stream, path).await
})
}
fn export_ml_bytes(
&self,
ctx: EngineContext,
bytes: async_channel::Sender<Result<Vec<u8>, Error>>,
config: MlExportConfig,
) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let stream = self.export_ml(context, config).await?;
tokio::spawn(pump_export_to_channel(stream, bytes));
Ok(())
})
}
fn import_file(&self, ctx: EngineContext, path: PathBuf) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let open = || async {
tokio::fs::File::open(&path)
.await
.map_err(|e| Error::internal(format!("Failed to open {}: {e}", path.display())))
};
if let Some(mut client) = self.compressed_upload_client() {
let stream = import_stream(context.clone(), open().await?);
match client.import_surql(stream).await {
Ok(_) => return Ok(()),
Err(status) if rejected_request_encoding(&status) => {}
Err(status) => return Err(status_to_error(status)),
}
}
self.client()
.import_surql(import_stream(context, open().await?))
.await
.map_err(status_to_error)?;
Ok(())
})
}
fn import_ml_file(&self, ctx: EngineContext, path: PathBuf) -> EngineFuture<'_, ()> {
Box::pin(async move {
let context = self.ready(ctx).await?;
let file = tokio::fs::File::open(&path)
.await
.map_err(|e| Error::internal(format!("Failed to open {}: {e}", path.display())))?;
self.client()
.import_ml_model(ml_import_stream(context, file))
.await
.map_err(status_to_error)?;
Ok(())
})
}
}
type ExportChunks =
std::pin::Pin<Box<dyn futures::Stream<Item = Result<tonic::codegen::Bytes, Error>> + Send>>;
impl GrpcEngine {
async fn export_surql(
&self,
context: rpc::RequestContext,
config: Option<DbExportConfig>,
) -> EngineResult<ExportChunks> {
let config = config.map(export_config).transpose()?;
let stream = self
.client()
.export_surql(self.checked(rpc::ExportSurqlRequest {
context: Some(context),
config,
})?)
.await
.map_err(status_to_error)?
.into_inner();
Ok(export_chunks(stream))
}
async fn export_ml(
&self,
context: rpc::RequestContext,
config: MlExportConfig,
) -> EngineResult<ExportChunks> {
let stream = self
.client()
.export_ml_model(self.checked(rpc::ExportMlModelRequest {
context: Some(context),
name: config.name,
version: config.version,
})?)
.await
.map_err(status_to_error)?
.into_inner();
Ok(export_chunks(stream))
}
}
enum ExportFrame {
Chunk(tonic::codegen::Bytes),
Trailer(rpc::DataTrailer),
Error(proto::SurrealError),
}
trait ExportResponse: Send + 'static {
fn into_frame(self) -> Option<ExportFrame>;
}
macro_rules! export_response {
($response:ty, $frame:path) => {
impl ExportResponse for $response {
fn into_frame(self) -> Option<ExportFrame> {
use $frame as Frame;
match self.frame? {
Frame::Chunk(chunk) => Some(ExportFrame::Chunk(chunk.data)),
Frame::Trailer(trailer) => Some(ExportFrame::Trailer(trailer)),
Frame::Error(error) => Some(ExportFrame::Error(error)),
}
}
}
};
}
export_response!(rpc::ExportSurqlResponse, rpc::export_surql_response::Frame);
export_response!(rpc::ExportMlModelResponse, rpc::export_ml_model_response::Frame);
fn export_chunks<T: ExportResponse>(stream: tonic::Streaming<T>) -> ExportChunks {
struct State<T> {
stream: tonic::Streaming<T>,
streamed: u64,
finished: bool,
}
impl<T> State<T> {
fn fail(mut self, error: Error) -> Option<(Result<tonic::codegen::Bytes, Error>, Self)> {
self.finished = true;
Some((Err(error), self))
}
fn truncated(
self,
message: String,
) -> Option<(Result<tonic::codegen::Bytes, Error>, Self)> {
self.fail(Error::connection(message, crate::types::ConnectionError::ConnectionFailed))
}
}
Box::pin(futures::stream::unfold(
State {
stream,
streamed: 0,
finished: false,
},
|mut state| async move {
if state.finished {
return None;
}
let frame = match state.stream.message().await {
Ok(Some(message)) => message.into_frame(),
Ok(None) => {
return state.truncated("The export ended before it was complete".to_string());
}
Err(status) => return state.fail(status_to_error(status)),
};
match frame {
Some(ExportFrame::Chunk(chunk)) => {
state.streamed += chunk.len() as u64;
Some((Ok(chunk), state))
}
Some(ExportFrame::Trailer(trailer)) if trailer.bytes == state.streamed => None,
Some(ExportFrame::Trailer(trailer)) => {
let message = format!(
"The export declared {} bytes but {} arrived",
trailer.bytes, state.streamed
);
state.truncated(message)
}
Some(ExportFrame::Error(error)) => state.fail(proto_error(error)),
None => state
.fail(Error::internal("The export carried an unrecognised frame".to_string())),
}
},
))
}
trait ImportRequest: Send + 'static {
fn chunk(data: tonic::codegen::Bytes) -> Self;
fn trailer(trailer: rpc::DataTrailer) -> Self;
}
impl ImportRequest for rpc::ImportSurqlRequest {
fn chunk(data: tonic::codegen::Bytes) -> Self {
Self {
frame: Some(rpc::import_surql_request::Frame::Chunk(rpc::DataChunk {
data,
})),
}
}
fn trailer(trailer: rpc::DataTrailer) -> Self {
Self {
frame: Some(rpc::import_surql_request::Frame::Trailer(trailer)),
}
}
}
impl ImportRequest for rpc::ImportMlModelRequest {
fn chunk(data: tonic::codegen::Bytes) -> Self {
Self {
frame: Some(rpc::import_ml_model_request::Frame::Chunk(rpc::DataChunk {
data,
})),
}
}
fn trailer(trailer: rpc::DataTrailer) -> Self {
Self {
frame: Some(rpc::import_ml_model_request::Frame::Trailer(trailer)),
}
}
}
fn import_stream(
context: rpc::RequestContext,
file: tokio::fs::File,
) -> impl futures::Stream<Item = rpc::ImportSurqlRequest> + Send + 'static {
let begin = rpc::ImportSurqlRequest {
frame: Some(rpc::import_surql_request::Frame::Begin(rpc::ImportSurqlBegin {
context: Some(context),
})),
};
frame_import(begin, file)
}
fn ml_import_stream(
context: rpc::RequestContext,
file: tokio::fs::File,
) -> impl futures::Stream<Item = rpc::ImportMlModelRequest> + Send + 'static {
let begin = rpc::ImportMlModelRequest {
frame: Some(rpc::import_ml_model_request::Frame::Begin(rpc::ImportMlModelBegin {
context: Some(context),
name: String::new(),
version: String::new(),
})),
};
frame_import(begin, file)
}
fn frame_import<T: ImportRequest>(
begin: T,
file: tokio::fs::File,
) -> impl futures::Stream<Item = T> + Send + 'static {
use tokio::io::AsyncReadExt;
enum Stage<T> {
Begin(T, tokio::fs::File),
Chunks(tokio::fs::File, u64),
Done,
}
futures::stream::unfold(Stage::Begin(begin, file), |stage| async move {
match stage {
Stage::Begin(begin, file) => Some((begin, Stage::Chunks(file, 0))),
Stage::Chunks(mut file, sent) => {
let mut buffer = Vec::with_capacity(surrealdb_protocol::DEFAULT_FILE_CHUNK_SIZE);
match file.read_buf(&mut buffer).await {
Ok(0) => Some((
T::trailer(rpc::DataTrailer {
bytes: sent,
blake3: String::new(),
}),
Stage::Done,
)),
Ok(read) => {
Some((T::chunk(buffer.into()), Stage::Chunks(file, sent + read as u64)))
}
Err(_) => None,
}
}
Stage::Done => None,
}
})
}
async fn write_export_to_file(mut stream: ExportChunks, path: PathBuf) -> EngineResult<()> {
use futures::StreamExt;
use tokio::io::AsyncWriteExt;
let mut file = tokio::fs::File::create(&path)
.await
.map_err(|e| Error::internal(format!("Failed to create {}: {e}", path.display())))?;
while let Some(chunk) = stream.next().await {
let chunk = chunk?;
file.write_all(&chunk)
.await
.map_err(|e| Error::internal(format!("Failed to write {}: {e}", path.display())))?;
}
file.flush()
.await
.map_err(|e| Error::internal(format!("Failed to flush {}: {e}", path.display())))?;
Ok(())
}
async fn pump_export_to_channel(
mut stream: ExportChunks,
bytes: async_channel::Sender<Result<Vec<u8>, Error>>,
) {
use futures::StreamExt;
while let Some(chunk) = stream.next().await {
let chunk = match chunk {
Ok(chunk) => chunk,
Err(error) => {
bytes.send(Err(error)).await.ok();
return;
}
};
if bytes.send(Ok(chunk.to_vec())).await.is_err() {
return;
}
}
bytes.close();
}
async fn pump_notifications(
mut stream: tonic::Streaming<rpc::SubscribeResponse>,
notifications: async_channel::Sender<Result<Notification, Error>>,
live_query_id: Uuid,
) {
loop {
let message = match stream.message().await {
Ok(Some(message)) => message,
Ok(None) => break,
Err(status) => {
notifications.send(Err(status_to_error(status))).await.ok();
break;
}
};
match message.frame {
Some(rpc::subscribe_response::Frame::Begin(_)) => {}
Some(rpc::subscribe_response::Frame::Notification(notification)) => {
match convert_notification(notification, live_query_id) {
Ok(notification) => {
if notifications.send(Ok(notification)).await.is_err() {
return;
}
}
Err(error) => {
notifications.send(Err(error)).await.ok();
return;
}
}
}
Some(rpc::subscribe_response::Frame::End(_)) => break,
Some(rpc::subscribe_response::Frame::Error(error)) => {
notifications.send(Err(proto_error(error))).await.ok();
break;
}
None => break,
}
}
notifications.close();
}
fn convert_notification(
notification: rpc::Notification,
live_query_id: Uuid,
) -> EngineResult<Notification> {
let action = match notification.action() {
rpc::Action::Created => Action::Create,
rpc::Action::Updated => Action::Update,
rpc::Action::Deleted => Action::Delete,
rpc::Action::Unspecified => {
return Err(Error::internal("Notification carried an unrecognised action".to_string()));
}
};
let record = match notification.record_id {
Some(record_id) => Value::RecordId(
record_id.try_into().map_err(|e: anyhow::Error| deserialization_error(e))?,
),
None => Value::None,
};
let result = match notification.value {
Some(value) => Value::try_from(value).map_err(deserialization_error)?,
None => Value::None,
};
let id = notification.live_query_id.and_then(|id| id.to_uuid().ok()).unwrap_or(live_query_id);
Ok(Notification::new(id.into(), None, action, record, result))
}
fn export_config(config: DbExportConfig) -> EngineResult<rpc::ExportConfig> {
use surrealdb_rpc::export::TableConfig;
if !config.database_definition {
return Err(Error::validation(
"Excluding the database definition from an export is not supported over gRPC: the \
protocol carries no field for it"
.to_string(),
ValidationError::InvalidRequest,
));
}
let tables = match config.tables {
TableConfig::All => rpc::export_config::Tables::from(true),
TableConfig::None => rpc::export_config::Tables::from(false),
TableConfig::Some(tables) => rpc::export_config::Tables {
selection: Some(rpc::export_config::tables::Selection::Selected(
rpc::export_config::SelectedTables {
tables,
},
)),
},
TableConfig::Exclude(excluded) => rpc::export_config::Tables {
selection: Some(rpc::export_config::tables::Selection::Excluded(
rpc::export_config::ExcludedTables {
tables: excluded.exclude,
},
)),
},
};
Ok(rpc::ExportConfig {
users: config.users,
accesses: config.accesses,
params: config.params,
functions: config.functions,
analyzers: config.analyzers,
tables: Some(tables),
versions: config.versions,
records: config.records,
sequences: config.sequences,
apis: config.apis,
buckets: config.buckets,
modules: config.modules,
configs: config.configs,
})
}
#[derive(Default)]
struct Statement {
values: Vec<Value>,
single: bool,
stats: Option<rpc::QueryStats>,
statement_kind: i32,
error: Option<Error>,
}
impl Statement {
fn push(&mut self, batch: rpc::QueryBatchFrame) {
self.statement_kind = batch.statement_kind;
if batch.kind == rpc::QueryResponseKind::Single as i32 {
self.single = true;
}
if batch.stats.is_some() {
self.stats = batch.stats;
}
if let Some(error) = batch.error {
self.error = Some(proto_error(error));
return;
}
match batch.payload {
Some(rpc::query_batch_frame::Payload::Values(values)) => {
self.values.reserve(values.values.len());
for value in values.values {
match Value::try_from(value) {
Ok(value) => self.values.push(value),
Err(e) => {
self.error = Some(deserialization_error(e));
return;
}
}
}
}
Some(rpc::query_batch_frame::Payload::Arrow(_)) => {
self.error = Some(Error::internal(
"Server sent a columnar batch, which was not requested".to_string(),
));
}
None => {}
}
}
fn finish(self) -> QueryResult {
let time = self
.stats
.and_then(|stats| stats.execution_duration)
.and_then(|duration| StdDuration::try_from(duration).ok())
.unwrap_or_default();
let query_type = if self.statement_kind == rpc::QueryStatementKind::Live as i32 {
QueryType::Live
} else if self.statement_kind == rpc::QueryStatementKind::Kill as i32 {
QueryType::Kill
} else {
QueryType::Other
};
let result = match self.error {
Some(error) => Err(error),
None if self.single => Ok(self.values.into_iter().next().unwrap_or(Value::None)),
None => Ok(Value::Array(Array::from(self.values))),
};
QueryResult {
time,
result,
query_type,
}
}
}
fn nullable(value: String) -> rpc::NullableString {
rpc::NullableString {
value: Some(rpc::nullable_string::Value::Some(value)),
}
}
fn to_proto_variables(variables: Variables) -> EngineResult<proto::Variables> {
to_variables(Object::from(variables).into_inner())
}
fn tokens_to_token(tokens: Option<rpc::Tokens>) -> EngineResult<Token> {
let tokens =
tokens.ok_or_else(|| Error::internal("Server did not return any tokens".to_string()))?;
Ok(if tokens.refresh.is_empty() {
Token::Access(tokens.access)
} else {
Token::WithRefresh {
access: tokens.access,
refresh: tokens.refresh,
}
})
}
fn access_method(credentials: Object) -> EngineResult<rpc::AccessMethod> {
let mut fields = credentials.into_inner();
let method = if let Some(access) = take_string(&mut fields, "ac")? {
let namespace = take_string(&mut fields, "ns")?.unwrap_or_default();
let database = take_string(&mut fields, "db")?.unwrap_or_default();
match take_string(&mut fields, "key")? {
Some(key) => rpc::access_method::Method::Bearer(rpc::BearerCredentials {
namespace,
database,
access,
key,
}),
None => rpc::access_method::Method::Record(rpc::RecordCredentials {
namespace,
database,
access,
variables: Some(to_variables(fields)?),
}),
}
} else if fields.contains_key("user") {
rpc::access_method::Method::User(rpc::UserCredentials {
namespace: take_string(&mut fields, "ns")?.unwrap_or_default(),
database: take_string(&mut fields, "db")?.unwrap_or_default(),
username: take_string(&mut fields, "user")?.unwrap_or_default(),
password: take_string(&mut fields, "pass")?.unwrap_or_default(),
access: String::new(),
})
} else {
return Err(Error::validation(
"Unrecognised credentials: expected a `user`, `ac`, or `key` field".to_string(),
None,
));
};
Ok(rpc::AccessMethod {
method: Some(method),
})
}
fn record_credentials(credentials: Object) -> EngineResult<rpc::RecordCredentials> {
let mut fields = credentials.into_inner();
let namespace = take_string(&mut fields, "ns")?.unwrap_or_default();
let database = take_string(&mut fields, "db")?.unwrap_or_default();
let access = take_string(&mut fields, "ac")?.ok_or_else(|| {
Error::validation("Missing `ac` field in signup credentials".to_string(), None)
})?;
Ok(rpc::RecordCredentials {
namespace,
database,
access,
variables: Some(to_variables(fields)?),
})
}
fn take_string(
fields: &mut std::collections::BTreeMap<String, Value>,
key: &str,
) -> EngineResult<Option<String>> {
match fields.remove(key) {
None => Ok(None),
Some(Value::String(value)) => Ok(Some(value)),
Some(_) => {
Err(Error::validation(format!("The `{key}` credential field must be a string"), None))
}
}
}
fn to_variables(
fields: std::collections::BTreeMap<String, Value>,
) -> EngineResult<proto::Variables> {
fields
.into_iter()
.map(|(key, value)| Ok((key, to_proto_value(value)?)))
.collect::<EngineResult<Vec<_>>>()
.map(|pairs| pairs.into_iter().collect())
}
fn to_proto_value(value: Value) -> EngineResult<proto::Value> {
proto::Value::try_from(value).map_err(|e| {
Error::serialization(e.to_string(), crate::types::SerializationError::Serialization)
})
}
#[allow(clippy::needless_pass_by_value)] fn status_to_error(status: tonic::Status) -> Error {
use tonic::Code;
let message = status.message().to_string();
match status.code() {
Code::InvalidArgument | Code::OutOfRange => Error::validation(message, None),
Code::Unimplemented => Error::configuration(message, None),
Code::DeadlineExceeded | Code::Aborted => Error::query(message, None),
Code::FailedPrecondition => {
if status.metadata().contains_key("surreal-rejected-before-execution") {
Error::validation(message, None)
} else {
Error::query(message, None)
}
}
Code::PermissionDenied | Code::Unauthenticated => Error::not_allowed(message, None),
Code::NotFound => Error::not_found(message, None),
Code::AlreadyExists => Error::already_exists(message, None),
Code::Unavailable | Code::Cancelled | Code::Unknown => Error::connection(message, None),
_ => Error::internal(message),
}
}
fn max_response_size(configured: Option<usize>, advertised: Option<usize>) -> Option<usize> {
configured.or(advertised)
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum LimitSource {
Server,
Client,
Both,
}
impl LimitSource {
fn remedy(self) -> &'static str {
match self {
Self::Server => "Raise SURREAL_GRPC_MAX_MESSAGE_SIZE on the server",
Self::Client => "Raise max_message_size on this connection's GrpcConfig",
Self::Both => {
"Raise SURREAL_GRPC_MAX_MESSAGE_SIZE on the server and max_message_size on \
this connection's GrpcConfig"
}
}
}
}
fn max_request_size(
configured: Option<usize>,
advertised: Option<usize>,
) -> Option<(usize, LimitSource)> {
match (configured, advertised) {
(Some(configured), Some(advertised)) => Some(match configured.cmp(&advertised) {
std::cmp::Ordering::Less => (configured, LimitSource::Client),
std::cmp::Ordering::Greater => (advertised, LimitSource::Server),
std::cmp::Ordering::Equal => (configured, LimitSource::Both),
}),
(Some(configured), None) => Some((configured, LimitSource::Client)),
(None, Some(advertised)) => Some((advertised, LimitSource::Server)),
(None, None) => None,
}
}
fn check_message_size<M: Message>(
limit: Option<(usize, LimitSource)>,
request: &M,
) -> EngineResult<()> {
let Some((limit, source)) = limit else {
return Ok(());
};
let size = request.encoded_len();
if size <= limit {
return Ok(());
}
Err(Error::validation(
format!(
"The request encodes to {size} bytes, above this connection's gRPC message limit of \
{limit} bytes. {}, or split the work into smaller statements.",
source.remedy()
),
ValidationError::InvalidRequest,
))
}
fn kind_str(kind: proto::ErrorKind) -> &'static str {
use proto::ErrorKind;
match kind {
ErrorKind::Validation => "Validation",
ErrorKind::Configuration => "Configuration",
ErrorKind::Query => "Query",
ErrorKind::Serialization => "Serialization",
ErrorKind::NotAllowed => "NotAllowed",
ErrorKind::NotFound => "NotFound",
ErrorKind::AlreadyExists => "AlreadyExists",
ErrorKind::Connection => "Connection",
ErrorKind::Thrown => "Thrown",
ErrorKind::Context => "Context",
ErrorKind::Internal | ErrorKind::Unspecified => "Internal",
}
}
fn proto_details(details: proto::ErrorDetails) -> Option<Value> {
if details.kind.is_empty() {
return None;
}
let mut object = Object::new();
object.insert("kind".to_string(), Value::String(details.kind));
if let Some(content) = details.content.and_then(|content| Value::try_from(content).ok()) {
object.insert("details".to_string(), content);
}
Some(Value::Object(object))
}
fn proto_error(error: proto::SurrealError) -> Error {
let kind = error.kind_or_internal();
let cause = error.cause.map(|cause| proto_error(*cause));
let details = error.details.and_then(proto_details);
let mapped = Error::from_wire(error.message, Some(kind_str(kind)), details);
match cause {
Some(cause) => mapped.with_cause(cause),
None => mapped,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{AuthError, NotAllowedError, Number, QueryError, SurrealValue};
fn object(pairs: &[(&str, &str)]) -> Object {
let mut object = Object::new();
for (key, value) in pairs {
object.insert((*key).to_string(), Value::String((*value).to_string()));
}
object
}
#[test]
fn user_credentials_are_classified_by_their_user_field() {
let method = access_method(object(&[("user", "root"), ("pass", "secret")]))
.expect("root credentials should classify");
let Some(rpc::access_method::Method::User(user)) = method.method else {
panic!("expected user credentials");
};
assert_eq!(user.username, "root");
assert_eq!(user.password, "secret");
assert_eq!(user.namespace, "");
assert_eq!(user.database, "");
}
#[test]
fn database_user_credentials_carry_their_scope() {
let method = access_method(object(&[
("ns", "test-ns"),
("db", "test-db"),
("user", "alice"),
("pass", "secret"),
]))
.expect("database credentials should classify");
let Some(rpc::access_method::Method::User(user)) = method.method else {
panic!("expected user credentials");
};
assert_eq!((user.namespace.as_str(), user.database.as_str()), ("test-ns", "test-db"));
}
#[test]
fn record_credentials_pass_their_extra_fields_as_variables() {
let method =
access_method(object(&[("ns", "n"), ("db", "d"), ("ac", "user"), ("email", "a@b.c")]))
.expect("record credentials should classify");
let Some(rpc::access_method::Method::Record(record)) = method.method else {
panic!("expected record credentials");
};
assert_eq!(record.access, "user");
let variables = record.variables.expect("variables should be set").variables;
assert_eq!(variables.len(), 1);
assert_eq!(variables[0].key, "email");
}
#[test]
fn bearer_credentials_are_classified_by_their_key_field() {
let method = access_method(object(&[("ac", "api"), ("key", "secret-key")]))
.expect("bearer credentials should classify");
assert!(matches!(method.method, Some(rpc::access_method::Method::Bearer(_))));
}
#[test]
fn record_credentials_keep_their_user_parameter() {
let method = access_method(object(&[
("ns", "n"),
("db", "d"),
("ac", "account"),
("user", "tobie"),
("tenant", "acme"),
]))
.expect("record credentials should classify");
let Some(rpc::access_method::Method::Record(record)) = method.method else {
panic!("expected record credentials");
};
assert_eq!(record.access, "account");
let variables = record.variables.expect("variables should be set").variables;
let keys: Vec<&str> = variables.iter().map(|kv| kv.key.as_str()).collect();
assert_eq!(keys, ["tenant", "user"]);
}
#[test]
fn non_string_credential_fields_are_rejected() {
let mut credentials = Object::new();
credentials.insert("user".to_string(), Value::String("root".to_string()));
credentials.insert("pass".to_string(), Value::Number(Number::Int(1234)));
let error =
access_method(credentials).expect_err("a non-string password should be rejected");
assert!(error.is_validation(), "expected a validation error, got {error:?}");
}
#[test]
fn the_replay_log_coalesces_repeated_operations() {
let entry = GrpcSession::new(Uuid::nil());
for _ in 0..10 {
entry.record(Replayable::Use {
namespace: Some("ns".to_string()),
database: Some("db".to_string()),
});
entry.record(Replayable::Set {
key: "tenant".to_string(),
value: Value::String("acme".to_string()),
});
}
let log = entry.log();
assert_eq!(log.len(), 2, "expected one `Use` and one `Set`, got {log:?}");
assert!(matches!(log[0], Replayable::Use { .. }));
assert!(matches!(&log[1], Replayable::Set { key, .. } if key == "tenant"));
}
#[test]
fn the_replay_log_does_not_coalesce_across_a_sign_in() {
let entry = GrpcSession::new(Uuid::nil());
let use_ns_db = || Replayable::Use {
namespace: Some("ns".to_string()),
database: Some("db".to_string()),
};
entry.record(use_ns_db());
entry.record(Replayable::Signin(Object::new()));
entry.record(use_ns_db());
let log = entry.log();
assert_eq!(log.len(), 3, "the `Use` before the sign-in must survive, got {log:?}");
}
#[test]
fn unrecognised_credentials_are_rejected() {
let error = access_method(object(&[("nonsense", "value")]))
.expect_err("unrecognised credentials should be rejected");
assert!(error.is_validation(), "expected a validation error, got {error:?}");
}
#[test]
fn signup_requires_an_access_method() {
let error = record_credentials(object(&[("ns", "n"), ("db", "d")]))
.expect_err("signup without `ac` should be rejected");
assert!(error.is_validation(), "expected a validation error, got {error:?}");
}
#[test]
fn wire_errors_keep_their_kind_and_cause() {
let inner = proto::SurrealError {
kind: proto::ErrorKind::NotFound as i32,
message: "no such table".to_string(),
..Default::default()
};
let outer = proto::SurrealError {
kind: proto::ErrorKind::Query as i32,
message: "query failed".to_string(),
cause: Some(Box::new(inner)),
..Default::default()
};
let error = proto_error(outer);
assert!(error.is_query(), "expected a query error, got {error:?}");
assert_eq!(error.message(), "query failed");
let cause = error.cause().expect("the cause should be preserved");
assert!(cause.is_not_found(), "expected a not-found cause, got {cause:?}");
}
#[test]
fn unknown_wire_error_kinds_degrade_to_internal() {
let error = proto_error(proto::SurrealError {
kind: 9999,
message: "from the future".to_string(),
..Default::default()
});
assert!(error.is_internal(), "expected an internal error, got {error:?}");
}
#[test]
fn wire_errors_rebuild_their_typed_reason() {
let conflict = proto_error(
proto::SurrealError::new(proto::ErrorKind::Query, "conflict")
.with_details("TransactionConflict", None),
);
assert_eq!(conflict.query_details(), Some(&QueryError::TransactionConflict));
let expired = proto_error(
proto::SurrealError::new(proto::ErrorKind::NotAllowed, "expired").with_details(
"Auth",
proto::Value::try_from(AuthError::TokenExpired.into_value()).ok(),
),
);
assert_eq!(
expired.not_allowed_details(),
Some(&NotAllowedError::Auth(AuthError::TokenExpired))
);
let detail = QueryError::TimedOut {
duration: StdDuration::from_secs(5),
};
let Value::Object(mut serialized) = detail.clone().into_value() else {
panic!("a detail serializes as a tagged object");
};
let content = serialized.remove("details").expect("TimedOut carries a duration");
let timed_out = proto_error(
proto::SurrealError::new(proto::ErrorKind::Query, "too slow")
.with_details("TimedOut", proto::Value::try_from(content).ok()),
);
assert_eq!(timed_out.query_details(), Some(&detail));
}
#[test]
fn a_retry_hint_does_not_invent_a_reason() {
let bare = proto_error(
proto::SurrealError::new(proto::ErrorKind::Query, "wording that will change")
.with_retry(None),
);
assert_eq!(bare.query_details(), None, "no reason was stated, so none is reported");
let timed_out = proto_error(
proto::SurrealError::new(proto::ErrorKind::Query, "too slow")
.with_details("TimedOut", None)
.with_retry(None),
);
assert!(timed_out.is_query(), "expected a query error, got {timed_out:?}");
let conflict = proto_error(
proto::SurrealError::new(proto::ErrorKind::Query, "conflict")
.with_details("TransactionConflict", None)
.with_retry(None),
);
assert_eq!(conflict.query_details(), Some(&QueryError::TransactionConflict));
}
#[test]
fn the_two_message_limits_reconcile_independently() {
assert_eq!(max_response_size(Some(128), Some(4)), Some(128));
assert_eq!(max_request_size(Some(128), Some(4)), Some((4, LimitSource::Server)));
assert_eq!(max_response_size(Some(2), Some(4)), Some(2));
assert_eq!(max_request_size(Some(2), Some(4)), Some((2, LimitSource::Client)));
assert_eq!(max_request_size(Some(4), Some(4)), Some((4, LimitSource::Both)));
assert_eq!(
(max_response_size(None, Some(4)), max_request_size(None, Some(4))),
(Some(4), Some((4, LimitSource::Server)))
);
assert_eq!(
(max_response_size(Some(8), None), max_request_size(Some(8), None)),
(Some(8), Some((8, LimitSource::Client)))
);
assert_eq!((max_response_size(None, None), max_request_size(None, None)), (None, None));
}
#[test]
fn an_oversized_request_is_refused_before_it_is_sent() {
let request = rpc::QueryRequest {
context: None,
query: "x".repeat(1024),
variables: None,
accepted_encodings: Vec::new(),
max_batch_records: 0,
};
let error = check_message_size(Some((64, LimitSource::Server)), &request)
.expect_err("above the limit");
let message = error.message();
assert!(message.contains("64 bytes"), "the limit should be named: {message}");
assert!(
message.contains(&request.encoded_len().to_string()),
"the size should be named: {message}"
);
check_message_size(Some((1 << 20, LimitSource::Server)), &request)
.expect("within the limit");
check_message_size(None, &request).expect("unbounded");
}
fn batch(values: Vec<Value>, kind: rpc::QueryResponseKind) -> rpc::QueryBatchFrame {
rpc::QueryBatchFrame {
query_index: 0,
batch_index: 0,
kind: kind as i32,
statement_kind: rpc::QueryStatementKind::Other as i32,
stats: None,
error: None,
payload: Some(rpc::query_batch_frame::Payload::Values(rpc::ValueBatch {
values: values
.into_iter()
.map(|v| proto::Value::try_from(v).expect("encodable"))
.collect(),
})),
}
}
#[test]
fn single_statements_are_not_wrapped_in_an_array() {
let mut statement = Statement::default();
statement.push(batch(vec![Value::Number(1.into())], rpc::QueryResponseKind::Single));
assert_eq!(statement.finish().result.unwrap(), Value::Number(1.into()));
let mut statement = Statement::default();
statement.push(batch(vec![Value::Number(1.into())], rpc::QueryResponseKind::BatchedFinal));
assert_eq!(
statement.finish().result.unwrap(),
Value::Array(Array::from(vec![Value::Number(1.into())]))
);
}
#[test]
fn batches_accumulate_across_frames() {
let mut statement = Statement::default();
statement.push(batch(vec![Value::Number(1.into())], rpc::QueryResponseKind::Batched));
statement.push(batch(vec![Value::Number(2.into())], rpc::QueryResponseKind::BatchedFinal));
assert_eq!(
statement.finish().result.unwrap(),
Value::Array(Array::from(vec![Value::Number(1.into()), Value::Number(2.into())]))
);
}
#[test]
fn indexes_that_never_emitted_are_not_results() {
let mut ran = Statement::default();
ran.push(batch(vec![Value::Number(1.into())], rpc::QueryResponseKind::BatchedFinal));
let mut statements: Vec<Option<Statement>> = Vec::new();
grow_statements(&mut statements, 3).expect("three statements is a legal count");
statements[1] = Some(ran);
let results = finish_statements(statements);
assert_eq!(results.len(), 1, "only the statement that emitted is a result");
assert_eq!(
results[0].result.as_ref().unwrap(),
&Value::Array(Array::from(vec![Value::Number(1.into())]))
);
}
#[test]
fn an_absurd_statement_count_is_refused() {
let mut statements: Vec<Option<Statement>> = Vec::new();
assert!(grow_statements(&mut statements, MAX_STATEMENTS + 1).is_err());
assert!(statements.is_empty(), "nothing is allocated for a refused count");
}
#[test]
fn excluding_tables_from_an_export_is_carried() {
use rpc::export_config::tables::Selection;
use surrealdb_rpc::export::{ExcludedTables, TableConfig};
let config = DbExportConfig {
tables: TableConfig::Exclude(ExcludedTables {
exclude: vec!["secrets".to_string()],
}),
..Default::default()
};
let wire = export_config(config).expect("excluding tables should be carried");
let Some(Selection::Excluded(excluded)) = wire.tables.and_then(|t| t.selection) else {
panic!("expected an excluded selection");
};
assert_eq!(excluded.tables, ["secrets"]);
}
#[test]
fn export_table_selection_maps_onto_the_wire() {
use rpc::export_config::tables::Selection;
use surrealdb_rpc::export::TableConfig;
let selection = |tables| {
export_config(DbExportConfig {
tables,
..Default::default()
})
.expect("selection should convert")
.tables
.expect("tables should be set")
.selection
.expect("a selection should be set")
};
assert!(matches!(selection(TableConfig::All), Selection::All(_)));
assert!(matches!(selection(TableConfig::None), Selection::None(_)));
let Selection::Selected(selected) =
selection(TableConfig::Some(vec!["person".to_string()]))
else {
panic!("expected a selected-tables selection");
};
assert_eq!(selected.tables, vec!["person".to_string()]);
}
#[test]
fn a_silent_server_is_assumed_to_support_everything() {
let features = extra_features(&rpc::ServerCapabilities::default());
assert!(features.contains(&ExtraFeatures::Backup));
assert!(features.contains(&ExtraFeatures::LiveQueries));
}
#[test]
fn live_queries_are_gated_on_the_reported_capability() {
let without = extra_features(&rpc::ServerCapabilities {
capabilities: vec!["TRANSACTIONS".to_string()],
..Default::default()
});
assert!(!without.contains(&ExtraFeatures::LiveQueries));
let with = extra_features(&rpc::ServerCapabilities {
capabilities: vec!["LIVE_QUERIES".to_string()],
..Default::default()
});
assert!(with.contains(&ExtraFeatures::LiveQueries));
}
#[test]
fn a_denied_export_method_withdraws_the_backup_feature() {
let features = extra_features(&rpc::ServerCapabilities {
capabilities: vec!["LIVE_QUERIES".to_string()],
denied_methods: vec![
"surrealdb.protocol.rpc.v1.SurrealDBService/ExportSurql".to_string(),
],
..Default::default()
});
assert!(!features.contains(&ExtraFeatures::Backup));
}
#[test]
fn an_export_deselecting_the_database_definition_is_refused() {
let refused = export_config(DbExportConfig {
database_definition: false,
..Default::default()
})
.expect_err("a selection this transport cannot carry must not be silently dropped");
assert!(
refused.is_validation(),
"the caller must be told to change the request: {refused:?}"
);
assert!(
refused.to_string().contains("database definition"),
"the refusal must say which selection could not be carried: {refused}"
);
export_config(DbExportConfig::default()).expect("the default selection must still cross");
}
#[test]
fn notifications_decode_their_record_and_value() {
let id = Uuid::from_u128(7);
let notification = convert_notification(
rpc::Notification {
live_query_id: Some(proto::Uuid::from_uuid(id)),
action: rpc::Action::Updated as i32,
record_id: None,
value: Some(
proto::Value::try_from(Value::String("hello".to_string())).expect("encodable"),
),
cursor: None,
},
id,
)
.expect("the notification should convert");
assert_eq!(notification.id, id.into());
assert_eq!(notification.action, Action::Update);
assert_eq!(notification.result, Value::String("hello".to_string()));
}
#[test]
fn notifications_with_an_unrecognised_action_are_rejected() {
let error = convert_notification(
rpc::Notification {
live_query_id: None,
action: rpc::Action::Unspecified as i32,
record_id: None,
value: None,
cursor: None,
},
Uuid::from_u128(1),
)
.expect_err("an unspecified action should be rejected");
assert!(error.is_internal(), "expected an internal error, got {error:?}");
}
}
#[cfg(test)]
mod request_compression_tests {
use super::rejected_request_encoding;
#[test]
fn an_unsupported_encoding_is_recognised() {
let mut status =
tonic::Status::unimplemented("Content is compressed with `zstd` which isn't supported");
status.metadata_mut().insert("grpc-accept-encoding", "identity".parse().unwrap());
assert!(rejected_request_encoding(&status));
}
#[test]
fn a_genuinely_unimplemented_method_is_not() {
let status = tonic::Status::unimplemented("Directory export is not served");
assert!(!rejected_request_encoding(&status));
}
#[test]
fn other_failures_are_not() {
for status in [
tonic::Status::unavailable("connection refused"),
tonic::Status::unauthenticated("bad credentials"),
tonic::Status::internal("boom"),
] {
assert!(!rejected_request_encoding(&status), "{status:?}");
}
}
}