use std::collections::{HashMap, VecDeque};
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use anyhow::Result;
use bytes::BytesMut;
use surrealdb_core::ctx::CancelHandle;
use surrealdb_core::dbs::Session;
use surrealdb_core::iam::verify::{self, ScramAuth, basic};
use surrealdb_core::kvs::{Datastore, Dialect, QueryRequest, QuerySource};
use surrealdb_core::syn;
use surrealdb_datastore::Transaction;
use surrealdb_kvs::TransactionType;
use surrealdb_rpc::capabilities::RouteTarget;
use surrealdb_sql::Ast;
use surrealdb_types::{Value, Variables};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpStream;
use tokio::sync::OwnedSemaphorePermit;
use tokio_rustls::TlsAcceptor;
use tokio_util::either::Either;
use tokio_util::sync::CancellationToken;
use super::error::PgError;
use super::msg::{self, DescribeTarget, Frontend, StartupMessage};
use super::typing::{PgColumn, PgType, ResultShape, shape_result, shape_result_jsonb};
use super::{CancelRegistry, LOG, encode, sasl};
use crate::cnf::PKG_VERSION;
type PgStream = Either<TcpStream, tokio_rustls::server::TlsStream<TcpStream>>;
const FORMAT_TEXT: i16 = 0;
const FORMAT_BINARY: i16 = 1;
const MAX_PREPARED: usize = 1024;
const MAX_PARAMS: usize = i16::MAX as usize;
struct PreparedStatement {
query: String,
parsed: Option<Ast>,
param_types: Vec<i32>,
param_count: usize,
empty: bool,
described: bool,
control: Option<Control>,
}
#[derive(Clone)]
enum Control {
Transaction(TxnOp),
Set(SetStatement),
}
impl Control {
fn command_tag(&self) -> &'static str {
match self {
Control::Transaction(TxnOp::Begin) => "BEGIN",
Control::Transaction(TxnOp::Commit) => "COMMIT",
Control::Transaction(TxnOp::Rollback) => "ROLLBACK",
Control::Set(_) => "SET",
}
}
}
struct Portal {
statement: String,
params: Variables,
result_formats: Vec<i16>,
executed: Option<Executed>,
control_ran: bool,
}
struct Executed {
columns: Vec<PgColumn>,
rows: VecDeque<Vec<Option<Value>>>,
}
const FLUSH_THRESHOLD: usize = 64 * 1024;
const STARTUP_TIMEOUT: Duration = Duration::from_secs(30);
const MAX_AUTH_MESSAGE_SIZE: usize = 4096;
#[allow(clippy::too_many_arguments)]
pub(super) async fn handle(
stream: TcpStream,
peer: SocketAddr,
ds: Arc<Datastore>,
ready: Arc<AtomicBool>,
shutdown: CancellationToken,
permit: OwnedSemaphorePermit,
registry: Arc<CancelRegistry>,
pid: i32,
secret: i32,
acceptor: Option<Arc<TlsAcceptor>>,
) {
let cancel = CancelHandle::new();
registry.insert((pid, secret), cancel.clone());
let negotiated = tokio::time::timeout(
STARTUP_TIMEOUT,
negotiate_tls(stream, acceptor.as_deref(), ®istry, &shutdown),
)
.await;
let Some((stream, prelude)) = negotiated.ok().flatten() else {
registry.remove(&(pid, secret));
return;
};
let mut conn = Connection {
stream,
prelude,
peer,
ds,
ready,
shutdown,
session: Session::default(),
statements: HashMap::new(),
portals: HashMap::new(),
dialect: Dialect::SurrealQl,
transaction: None,
txn_failed: false,
cancel,
registry: Arc::clone(®istry),
pid,
secret,
_permit: permit,
};
if let Err(err) = conn.run().await {
debug!(target: LOG, "postgres connection from {peer} ended: {err}");
}
if let Some(tx) = conn.transaction.take() {
let _ = tx.cancel().await;
}
registry.remove(&(pid, secret));
}
async fn negotiate_tls(
mut tcp: TcpStream,
acceptor: Option<&TlsAcceptor>,
registry: &CancelRegistry,
shutdown: &CancellationToken,
) -> Option<(PgStream, Option<Vec<u8>>)> {
loop {
let payload = read_startup_packet_raw(&mut tcp, shutdown).await?;
match msg::parse_startup(&payload) {
Ok(StartupMessage::SslRequest) => match acceptor {
Some(acceptor) => {
tcp.write_all(b"S").await.ok()?;
match acceptor.accept(tcp).await {
Ok(tls) => return Some((Either::Right(tls), None)),
Err(err) => {
debug!(target: LOG, "postgres TLS handshake failed: {err}");
return None;
}
}
}
None => tcp.write_all(b"N").await.ok()?,
},
Ok(StartupMessage::GssEncRequest) => tcp.write_all(b"N").await.ok()?,
Ok(StartupMessage::CancelRequest {
pid,
secret,
}) => {
if let Some(handle) = registry.get(&(pid, secret)) {
handle.trip();
}
return None;
}
Ok(StartupMessage::Startup {
..
}) => return Some((Either::Left(tcp), Some(payload))),
Err(_) => return None,
}
}
}
async fn read_startup_packet_raw(
tcp: &mut TcpStream,
shutdown: &CancellationToken,
) -> Option<Vec<u8>> {
let len = tokio::select! {
biased;
_ = shutdown.cancelled() => return None,
len = tcp.read_i32() => len.ok()?,
};
read_startup_payload(tcp, len).await.ok().flatten()
}
async fn read_startup_payload<R>(reader: &mut R, len: i32) -> Result<Option<Vec<u8>>>
where
R: AsyncReadExt + Unpin,
{
let len = usize::try_from(len).unwrap_or(0);
if !(8..=msg::MAX_STARTUP_PACKET_SIZE).contains(&len) {
return Ok(None);
}
let mut payload = vec![0u8; len - 4];
reader.read_exact(&mut payload).await?;
Ok(Some(payload))
}
pub(super) async fn reject_overloaded(mut stream: TcpStream) {
let _ = reject_overloaded_inner(&mut stream).await;
}
async fn reject_overloaded_inner(stream: &mut TcpStream) -> Result<()> {
loop {
let len = stream.read_i32().await?;
let Some(payload) = read_startup_payload(stream, len).await? else {
return Ok(());
};
match msg::parse_startup(&payload) {
Ok(StartupMessage::SslRequest) | Ok(StartupMessage::GssEncRequest) => {
let mut buf = BytesMut::new();
msg::write_ssl_response(&mut buf, false);
stream.write_all(&buf).await?;
}
_ => break,
}
}
let mut buf = BytesMut::new();
msg::write_error_response(&mut buf, &PgError::too_many_connections());
stream.write_all(&buf).await?;
stream.flush().await?;
Ok(())
}
struct Connection {
stream: PgStream,
prelude: Option<Vec<u8>>,
peer: SocketAddr,
ds: Arc<Datastore>,
ready: Arc<AtomicBool>,
shutdown: CancellationToken,
session: Session,
statements: HashMap<String, PreparedStatement>,
portals: HashMap<String, Portal>,
dialect: Dialect,
transaction: Option<Arc<Transaction>>,
txn_failed: bool,
cancel: CancelHandle,
registry: Arc<CancelRegistry>,
pid: i32,
secret: i32,
_permit: OwnedSemaphorePermit,
}
impl Connection {
async fn run(&mut self) -> Result<()> {
let authed = match tokio::time::timeout(STARTUP_TIMEOUT, self.startup_and_auth()).await {
Ok(res) => res?,
Err(_) => {
let _ = self.report(&PgError::auth_timeout()).await;
return Ok(());
}
};
if !authed {
return Ok(());
}
self.command_loop().await
}
async fn startup_and_auth(&mut self) -> Result<bool> {
let Some(params) = self.startup().await? else {
return Ok(false);
};
self.authenticate(params).await
}
async fn startup(&mut self) -> Result<Option<Vec<(String, String)>>> {
let payload = match self.prelude.take() {
Some(payload) => payload,
None => match self.read_startup_packet().await? {
Some(payload) => payload,
None => return Ok(None),
},
};
match msg::parse_startup(&payload) {
Ok(StartupMessage::Startup {
version,
params,
}) => {
let major = version >> 16;
let minor = version & 0xffff;
if major != 3 {
self.report(
&PgError::protocol(format!(
"unsupported protocol version {major}.{minor} (only 3.x is supported)"
))
.fatal(),
)
.await?;
return Ok(None);
}
if minor > 0 || params.iter().any(|(k, _)| k.starts_with("_pq_.")) {
let unsupported: Vec<String> = params
.iter()
.filter(|(k, _)| k.starts_with("_pq_."))
.map(|(k, _)| k.clone())
.collect();
let mut buf = BytesMut::new();
msg::write_negotiate_protocol_version(&mut buf, &unsupported);
self.stream.write_all(&buf).await?;
}
Ok(Some(params))
}
Ok(_) => {
self.report(&PgError::protocol("unexpected request during startup").fatal())
.await?;
Ok(None)
}
Err(err) => {
self.report(&err).await?;
Ok(None)
}
}
}
async fn authenticate(&mut self, params: Vec<(String, String)>) -> Result<bool> {
if !self.ds.allows_http_route(&RouteTarget::Postgres) {
warn!(
target: LOG,
"Capabilities denied postgres connection attempt from {}", self.peer
);
self.report(
&PgError::insufficient_privilege(
"Forbidden: the postgres protocol is not allowed on this server",
)
.fatal(),
)
.await?;
return Ok(false);
}
if !self.ready.load(Ordering::SeqCst) {
self.report(&PgError::cannot_connect_now("the database system is starting up").fatal())
.await?;
return Ok(false);
}
let mut user = None;
let mut database = None;
let mut options = None;
for (key, value) in params {
match key.as_str() {
"user" => user = Some(value),
"database" => database = Some(value),
"options" => options = Some(value),
_ => {}
}
}
let Some(user) = user else {
self.report(&PgError::protocol("no user specified in startup packet").fatal()).await?;
return Ok(false);
};
if let Some(options) = options.as_deref()
&& let Some(dialect) = dialect_from_options(options)
&& let Err(err) = self.select_dialect(dialect)
{
self.report(&err.fatal()).await?;
return Ok(false);
}
let (ns, db) = parse_database_param(&user, database.as_deref());
self.session = Session::default();
self.session.ip = Some(self.peer.ip().to_string());
self.session.id = Some(uuid::Uuid::new_v4());
if self.ds.is_auth_enabled() {
let scram =
match verify::scram_lookup(&self.ds, &user, ns.as_deref(), db.as_deref()).await {
Ok(scram) => scram,
Err(err) => {
debug!(target: LOG, "postgres SCRAM lookup failed for user '{user}': {err}");
self.report(&auth_failed(&user)).await?;
return Ok(false);
}
};
let authenticated = match scram {
Some(scram) => self.authenticate_scram(&user, &scram).await?,
None => self.authenticate_cleartext(&user, ns.as_deref(), db.as_deref()).await?,
};
if !authenticated {
return Ok(false);
}
}
if !self.ds.allows_query_by_subject(self.session.au.as_ref()) {
self.report(
&PgError::insufficient_privilege(
"Forbidden: this user is not allowed to query the database",
)
.fatal(),
)
.await?;
return Ok(false);
}
if ns.is_some() || db.is_some() {
if let Err(err) = self.ds.process_use(None, &mut self.session, ns, db).await {
self.report(&PgError::from(&err).fatal()).await?;
return Ok(false);
}
}
let mut buf = BytesMut::new();
msg::write_authentication_ok(&mut buf);
let server_version = format!("16.0-surrealdb.{}", *PKG_VERSION);
for (key, value) in [
("server_version", server_version.as_str()),
("server_encoding", "UTF8"),
("client_encoding", "UTF8"),
("DateStyle", "ISO, MDY"),
("integer_datetimes", "on"),
("standard_conforming_strings", "on"),
("TimeZone", "UTC"),
] {
msg::write_parameter_status(&mut buf, key, value);
}
msg::write_backend_key_data(&mut buf, self.pid, self.secret);
msg::write_ready_for_query(&mut buf, b'I');
self.stream.write_all(&buf).await?;
self.stream.flush().await?;
debug!(target: LOG, "postgres connection established from {} for user '{user}'", self.peer);
Ok(true)
}
async fn verify_credentials(
&mut self,
user: &str,
pass: &str,
ns: Option<&str>,
db: Option<&str>,
) -> bool {
if let (Some(ns), Some(db)) = (ns, db)
&& basic(&self.ds, &mut self.session, user, pass, Some(ns), Some(db)).await.is_ok()
{
return true;
}
if let Some(ns) = ns
&& basic(&self.ds, &mut self.session, user, pass, Some(ns), None).await.is_ok()
{
return true;
}
match basic(&self.ds, &mut self.session, user, pass, None, None).await {
Ok(()) => true,
Err(err) => {
debug!(target: LOG, "postgres authentication failed for user '{user}': {err}");
false
}
}
}
async fn authenticate_cleartext(
&mut self,
user: &str,
ns: Option<&str>,
db: Option<&str>,
) -> Result<bool> {
if !self.is_tls() {
self.report(
&PgError::invalid_password(
"password authentication requires an encrypted connection for this user; \
connect with TLS (sslmode=require) or define the user with a password so \
SCRAM material is generated",
)
.fatal(),
)
.await?;
return Ok(false);
}
let mut buf = BytesMut::new();
msg::write_authentication_cleartext_password(&mut buf);
self.stream.write_all(&buf).await?;
self.stream.flush().await?;
let Some((tag, payload)) = self.read_message(MAX_AUTH_MESSAGE_SIZE).await? else {
return Ok(false);
};
let Frontend::Password(pass) = msg::parse_frontend(tag, &payload)? else {
self.report(&PgError::protocol("expected a password message").fatal()).await?;
return Ok(false);
};
if !self.verify_credentials(user, &pass, ns, db).await {
self.report(&auth_failed(user)).await?;
return Ok(false);
}
Ok(true)
}
async fn authenticate_scram(&mut self, user: &str, scram: &ScramAuth) -> Result<bool> {
let mut buf = BytesMut::new();
msg::write_authentication_sasl(&mut buf, &[sasl::MECHANISM]);
self.stream.write_all(&buf).await?;
self.stream.flush().await?;
let Some((tag, payload)) = self.read_message(MAX_AUTH_MESSAGE_SIZE).await? else {
return Ok(false);
};
if tag != b'p' {
self.report(&PgError::protocol("expected a SASL response message").fatal()).await?;
return Ok(false);
}
let (mechanism, client_first) = match msg::parse_sasl_initial(&payload) {
Ok(parsed) => parsed,
Err(err) => {
self.report(&err.fatal()).await?;
return Ok(false);
}
};
if mechanism != sasl::MECHANISM {
self.report(
&PgError::protocol(format!("unsupported SASL mechanism \"{mechanism}\"")).fatal(),
)
.await?;
return Ok(false);
}
let (exchange, server_first) =
match sasl::ScramExchange::start(&client_first, scram.salt(), scram.iterations()) {
Ok(started) => started,
Err(err) => {
self.report(&err.fatal()).await?;
return Ok(false);
}
};
let mut buf = BytesMut::new();
msg::write_authentication_sasl_continue(&mut buf, server_first.as_bytes());
self.stream.write_all(&buf).await?;
self.stream.flush().await?;
let Some((tag, payload)) = self.read_message(MAX_AUTH_MESSAGE_SIZE).await? else {
return Ok(false);
};
if tag != b'p' {
self.report(&PgError::protocol("expected a SASL response message").fatal()).await?;
return Ok(false);
}
let server_final = match exchange.finish(&payload, scram) {
Ok(final_message) => final_message,
Err(_) => {
self.report(&auth_failed(user)).await?;
return Ok(false);
}
};
let mut buf = BytesMut::new();
msg::write_authentication_sasl_final(&mut buf, server_final.as_bytes());
self.stream.write_all(&buf).await?;
self.stream.flush().await?;
if let Err(err) = scram.apply(&mut self.session) {
self.report(&PgError::internal(err.to_string()).fatal()).await?;
return Ok(false);
}
Ok(true)
}
async fn command_loop(&mut self) -> Result<()> {
let mut out = BytesMut::new();
let mut skip_until_sync = false;
loop {
let Some((tag, payload)) = self.read_message(msg::MAX_MESSAGE_SIZE).await? else {
return Ok(());
};
let message = match msg::parse_frontend(tag, &payload) {
Ok(message) => message,
Err(err) => {
msg::write_error_response(&mut out, &err);
skip_until_sync = true;
continue;
}
};
match message {
Frontend::Query(sql) => {
self.flush(&mut out).await?;
skip_until_sync = false;
self.simple_query(&sql).await?;
}
Frontend::Parse {
name,
query,
param_types,
} => {
if skip_until_sync {
continue;
}
if transaction_control(&query).is_none()
&& let Some(err) = self.aborted_txn_guard()
{
msg::write_error_response(&mut out, &err);
skip_until_sync = true;
continue;
}
if let Err(err) = self.handle_parse(name, &query, param_types, &mut out) {
self.note_extended_error(&err, &mut out);
skip_until_sync = true;
}
}
Frontend::Bind {
portal,
statement,
param_formats,
params,
result_formats,
} => {
if skip_until_sync {
continue;
}
if !self.statement_is_txn_control(&statement)
&& let Some(err) = self.aborted_txn_guard()
{
msg::write_error_response(&mut out, &err);
skip_until_sync = true;
continue;
}
if let Err(err) = self.handle_bind(
portal,
statement,
¶m_formats,
¶ms,
result_formats,
&mut out,
) {
self.note_extended_error(&err, &mut out);
skip_until_sync = true;
}
}
Frontend::Describe {
target,
name,
} => {
if skip_until_sync {
continue;
}
let is_txn_control = match target {
DescribeTarget::Statement => self.statement_is_txn_control(&name),
DescribeTarget::Portal => self.portal_is_txn_control(&name),
};
if !is_txn_control && let Some(err) = self.aborted_txn_guard() {
msg::write_error_response(&mut out, &err);
skip_until_sync = true;
continue;
}
if let Err(err) = self.handle_describe(target, &name, &mut out).await {
self.note_extended_error(&err, &mut out);
skip_until_sync = true;
}
}
Frontend::Execute {
portal,
max_rows,
} => {
if skip_until_sync {
continue;
}
if !self.portal_is_txn_control(&portal)
&& let Some(err) = self.aborted_txn_guard()
{
msg::write_error_response(&mut out, &err);
skip_until_sync = true;
continue;
}
if let Err(err) = self.handle_execute(&portal, max_rows, &mut out).await? {
self.note_extended_error(&err, &mut out);
skip_until_sync = true;
}
}
Frontend::Close {
target,
name,
} => {
if skip_until_sync {
continue;
}
self.handle_close(target, &name, &mut out);
}
Frontend::Sync => {
skip_until_sync = false;
self.portals.clear();
let status = self.ready_status();
msg::write_ready_for_query(&mut out, status);
self.flush(&mut out).await?;
}
Frontend::Flush => self.flush(&mut out).await?,
Frontend::Terminate => return Ok(()),
Frontend::Password(_) => {
self.report(&PgError::protocol("unexpected password message").fatal()).await?;
return Ok(());
}
Frontend::Unknown(tag) => {
self.report(
&PgError::protocol(format!(
"unexpected message type '{}'",
char::from(tag)
))
.fatal(),
)
.await?;
return Ok(());
}
}
}
}
async fn flush(&mut self, out: &mut BytesMut) -> Result<()> {
if !out.is_empty() {
self.stream.write_all(out).await?;
out.clear();
}
self.stream.flush().await?;
Ok(())
}
fn handle_parse(
&mut self,
name: String,
query: &str,
param_types: Vec<i32>,
out: &mut BytesMut,
) -> Result<(), PgError> {
if !name.is_empty() && self.statements.contains_key(&name) {
return Err(PgError::duplicate_statement(&name));
}
if !self.statements.contains_key(&name) && self.statements.len() >= MAX_PREPARED {
return Err(PgError::too_many_prepared("prepared statements"));
}
if let Some(control) = transaction_control(query).map(Control::Transaction).or_else(|| {
parse_set(query).filter(|set| self.should_intercept_set(set)).map(Control::Set)
}) {
self.statements.insert(
name,
PreparedStatement {
query: query.to_string(),
parsed: None,
param_types: Vec::new(),
param_count: 0,
empty: false,
described: false,
control: Some(control),
},
);
msg::write_parse_complete(out);
return Ok(());
}
let (query, parsed, param_count, empty) = match self.dialect {
Dialect::SurrealQl => {
let (query, positional) = rewrite_positional_params(query);
let config = self.ds.parser_config();
let ast =
syn::parse_with_capabilities(&query, &self.ds.get_capabilities(), &config)
.map_err(|e| PgError::syntax(e.to_string()))?;
let empty = ast.num_statements() == 0;
let param_count = param_types.len().max(positional);
let parsed = if empty {
None
} else {
Some(ast)
};
(query, parsed, param_count, empty)
}
Dialect::Gql => {
let empty = query.trim().is_empty();
(query.to_string(), None, param_types.len(), empty)
}
};
if param_count > MAX_PARAMS {
return Err(PgError::protocol(format!(
"statement references {param_count} parameters, exceeding the maximum of {MAX_PARAMS}"
)));
}
self.statements.insert(
name,
PreparedStatement {
query,
parsed,
param_count,
param_types,
empty,
described: false,
control: None,
},
);
msg::write_parse_complete(out);
Ok(())
}
fn handle_bind(
&mut self,
portal: String,
statement: String,
param_formats: &[i16],
params: &[Option<Vec<u8>>],
result_formats: Vec<i16>,
out: &mut BytesMut,
) -> Result<(), PgError> {
let Some(stmt) = self.statements.get(&statement) else {
return Err(PgError::invalid_statement(&statement));
};
if params.len() < stmt.param_count {
return Err(PgError::protocol(format!(
"bind message supplies {} parameters, but prepared statement \"{statement}\" requires {}",
params.len(),
stmt.param_count
)));
}
if !self.portals.contains_key(&portal) && self.portals.len() >= MAX_PREPARED {
return Err(PgError::too_many_prepared("portals"));
}
let mut vars = Variables::new();
for (i, raw) in params.iter().enumerate().take(stmt.param_count) {
let oid = stmt.param_types.get(i).copied().unwrap_or(0);
let format = format_at(param_formats, i);
let value = encode::decode_param(raw.as_deref(), oid, format)?;
vars.insert(format!("_{}", i + 1), value);
}
self.portals.insert(
portal,
Portal {
statement,
params: vars,
result_formats,
executed: None,
control_ran: false,
},
);
msg::write_bind_complete(out);
Ok(())
}
async fn handle_describe(
&mut self,
target: DescribeTarget,
name: &str,
out: &mut BytesMut,
) -> Result<(), PgError> {
match target {
DescribeTarget::Statement => {
let Some(stmt) = self.statements.get_mut(name) else {
return Err(PgError::invalid_statement(name));
};
if stmt.control.is_some() {
msg::write_parameter_description(out, &[]);
msg::write_no_data(out);
return Ok(());
}
stmt.described = true;
let empty = stmt.empty;
let oids: Vec<i32> = (0..stmt.param_count)
.map(|i| stmt.param_types.get(i).copied().unwrap_or(0))
.collect();
msg::write_parameter_description(out, &oids);
if empty {
msg::write_no_data(out);
} else {
let columns = [PgColumn {
name: "result".to_string(),
ty: PgType::Jsonb,
}];
msg::write_row_description(out, &columns, |_| FORMAT_TEXT);
}
}
DescribeTarget::Portal => {
if self.portal_is_control(name) {
msg::write_no_data(out);
return Ok(());
}
if self.portal_is_empty(name)? {
msg::write_no_data(out);
return Ok(());
}
self.ensure_executed(name).await?;
let portal = self.portals.get(name).expect("ensured above");
let formats = portal.result_formats.clone();
let columns = &portal.executed.as_ref().expect("ensured above").columns;
msg::write_row_description(out, columns, |i| format_at(&formats, i));
}
}
Ok(())
}
async fn handle_execute(
&mut self,
portal_name: &str,
max_rows: i32,
out: &mut BytesMut,
) -> Result<Result<(), PgError>> {
if let Some(control) = self.portal_control(portal_name) {
if self.portals.get(portal_name).is_some_and(|p| p.control_ran) {
msg::write_command_complete(out, control.command_tag());
return Ok(Ok(()));
}
if let Some(portal) = self.portals.get_mut(portal_name) {
portal.control_ran = true;
}
let result = match control {
Control::Transaction(op) => self.handle_transaction_control(op, out).await,
Control::Set(set) => self.handle_set(set, out),
};
return Ok(result);
}
match self.portal_is_empty(portal_name) {
Ok(true) => {
msg::write_empty_query_response(out);
return Ok(Ok(()));
}
Ok(false) => {}
Err(err) => return Ok(Err(err)),
}
if let Err(err) = self.ensure_executed(portal_name).await {
return Ok(Err(err));
}
let (columns, formats) = {
let portal = self.portals.get(portal_name).expect("ensured above");
let executed = portal.executed.as_ref().expect("ensured above");
(executed.columns.clone(), portal.result_formats.clone())
};
let limit = if max_rows <= 0 {
usize::MAX
} else {
max_rows as usize
};
let mut sent = 0usize;
while sent < limit {
let row = {
let portal = self.portals.get_mut(portal_name).expect("ensured above");
portal.executed.as_mut().expect("ensured above").rows.pop_front()
};
let Some(row) = row else {
break;
};
let cells = match encode_row(row, &columns, &formats) {
Ok(cells) => cells,
Err(err) => return Ok(Err(err)),
};
msg::write_data_row(out, &cells);
sent += 1;
if out.len() >= FLUSH_THRESHOLD {
self.stream.write_all(out).await?;
out.clear();
}
}
let remaining = {
let portal = self.portals.get(portal_name).expect("ensured above");
!portal.executed.as_ref().expect("ensured above").rows.is_empty()
};
if remaining {
msg::write_portal_suspended(out);
} else {
msg::write_command_complete(out, &format!("SELECT {sent}"));
}
Ok(Ok(()))
}
fn portal_is_empty(&self, portal_name: &str) -> Result<bool, PgError> {
let portal =
self.portals.get(portal_name).ok_or_else(|| PgError::invalid_cursor(portal_name))?;
let stmt = self
.statements
.get(&portal.statement)
.ok_or_else(|| PgError::invalid_statement(&portal.statement))?;
Ok(stmt.empty)
}
fn statement_control(&self, name: &str) -> Option<&Control> {
self.statements.get(name)?.control.as_ref()
}
fn portal_control_ref(&self, portal_name: &str) -> Option<&Control> {
self.statement_control(&self.portals.get(portal_name)?.statement)
}
fn portal_control(&self, portal_name: &str) -> Option<Control> {
self.portal_control_ref(portal_name).cloned()
}
fn portal_is_control(&self, portal_name: &str) -> bool {
self.portal_control_ref(portal_name).is_some()
}
fn statement_is_txn_control(&self, name: &str) -> bool {
matches!(self.statement_control(name), Some(Control::Transaction(_)))
}
fn portal_is_txn_control(&self, portal_name: &str) -> bool {
matches!(self.portal_control_ref(portal_name), Some(Control::Transaction(_)))
}
fn handle_close(&mut self, target: DescribeTarget, name: &str, out: &mut BytesMut) {
match target {
DescribeTarget::Statement => {
self.statements.remove(name);
self.portals.retain(|_, p| p.statement != name);
}
DescribeTarget::Portal => {
self.portals.remove(name);
}
}
msg::write_close_complete(out);
}
async fn ensure_executed(&mut self, portal_name: &str) -> Result<(), PgError> {
if self.portals.get(portal_name).is_some_and(|p| p.executed.is_some()) {
return Ok(());
}
let Some(portal) = self.portals.get(portal_name) else {
return Err(PgError::invalid_cursor(portal_name));
};
let statement = portal.statement.clone();
let params = portal.params.clone();
let Some(stmt) = self.statements.get(&statement) else {
return Err(PgError::invalid_statement(&statement));
};
let jsonb_shape = stmt.described;
let parsed = stmt.parsed.clone();
let query = stmt.query.clone();
let value = self.execute_query(parsed, &query, params).await?;
let shape = if jsonb_shape {
shape_result_jsonb(value)
} else {
shape_result(value)
};
if shape.columns.len() > msg::MAX_COLUMNS {
return Err(PgError::feature_not_supported(format!(
"result has {} columns, exceeding the maximum of {}",
shape.columns.len(),
msg::MAX_COLUMNS
)));
}
let portal = self.portals.get_mut(portal_name).expect("checked above");
portal.executed = Some(Executed {
columns: shape.columns,
rows: shape.rows.into(),
});
Ok(())
}
async fn execute_query(
&mut self,
parsed: Option<Ast>,
query: &str,
params: Variables,
) -> Result<Value, PgError> {
self.check_current_dialect()?;
if self.dialect == Dialect::Gql && self.transaction.is_some() {
return Err(PgError::feature_not_supported(
"GQL queries cannot run inside an interactive transaction",
));
}
let cancel = self.arm_cancel();
let results = match (self.dialect, self.transaction.clone(), parsed) {
(Dialect::Gql, _, _) => self
.ds
.run(
QueryRequest::new(QuerySource::gql(query), &self.session)
.with_variables(Some(params))
.with_cancel(cancel),
)
.await
.map_err(|e| PgError::from(&e)),
(Dialect::SurrealQl, transaction, Some(ast)) => self
.ds
.run(
QueryRequest::new(ast, &self.session)
.with_variables(Some(params))
.with_optional_transaction(transaction)
.with_cancel(cancel),
)
.await
.map_err(|e| PgError::from(&e)),
(Dialect::SurrealQl, _, None) => Ok(Vec::new()),
}?;
let mut last = Value::None;
for result in results {
match result.result {
Ok(value) => last = value,
Err(err) => return Err(PgError::from(&err)),
}
}
Ok(last)
}
async fn simple_query(&mut self, sql: &str) -> Result<()> {
let mut out = BytesMut::new();
let trimmed = sql.trim();
if trimmed.is_empty() {
msg::write_empty_query_response(&mut out);
return self.finish_query(out).await;
}
if let Some(op) = transaction_control(trimmed) {
if let Err(err) = self.handle_transaction_control(op, &mut out).await {
msg::write_error_response(&mut out, &err);
}
return self.finish_query(out).await;
}
if self.txn_failed {
msg::write_error_response(
&mut out,
&PgError::in_failed_transaction(
"current transaction is aborted, commands ignored until end of transaction block",
),
);
return self.finish_query(out).await;
}
if let Some(set) = parse_set(trimmed).filter(|set| self.should_intercept_set(set)) {
if let Err(err) = self.handle_set(set, &mut out) {
self.note_simple_error(&err, &mut out);
}
return self.finish_query(out).await;
}
if let Err(err) = self.check_current_dialect() {
self.note_simple_error(&err, &mut out);
return self.finish_query(out).await;
}
match self.dialect {
Dialect::Gql if self.transaction.is_some() => {
self.note_simple_error(
&PgError::feature_not_supported(
"GQL queries cannot run inside an interactive transaction",
),
&mut out,
);
}
Dialect::Gql => self.run_gql(sql, &mut out).await?,
Dialect::SurrealQl if self.transaction.is_some() => {
self.run_in_transaction(sql, &mut out).await?;
}
Dialect::SurrealQl => self.run_surrealql_autocommit(sql, &mut out).await?,
}
self.finish_query(out).await
}
async fn run_surrealql_autocommit(&mut self, sql: &str, out: &mut BytesMut) -> Result<()> {
let config = self.ds.parser_config();
let ast = match syn::parse_with_capabilities(sql, &self.ds.get_capabilities(), &config) {
Ok(ast) => ast,
Err(err) => {
msg::write_error_response(out, &PgError::syntax(err.to_string()));
return Ok(());
}
};
if ast.num_statements() == 0 {
msg::write_empty_query_response(out);
return Ok(());
}
for mut unit in ast.into_execution_units() {
if unit.is_sole_commit() {
msg::write_notice(out, "25P01", "there is no transaction in progress");
msg::write_command_complete(out, "COMMIT");
continue;
}
if unit.is_sole_cancel() {
msg::write_notice(out, "25P01", "there is no transaction in progress");
msg::write_command_complete(out, "ROLLBACK");
continue;
}
let capture_session = !unit.contains_cancel();
let let_vars = if capture_session {
let names = unit.get_let_statements();
for name in &names {
unit.add_param(name.clone());
}
unit.add_param("session".to_string());
names
} else {
Vec::new()
};
let cancel = self.arm_cancel();
let mut results =
match self.ds.run(QueryRequest::new(unit, &self.session).with_cancel(cancel)).await
{
Ok(results) => results,
Err(err) => {
msg::write_error_response(out, &PgError::from(&err));
break;
}
};
let trailing = if capture_session && results.len() > let_vars.len() {
results.split_off(results.len() - (let_vars.len() + 1))
} else {
Vec::new()
};
let mut ok = true;
for result in results {
match result.result {
Ok(value) => {
if !self.emit_result(value, out).await? {
ok = false;
break;
}
}
Err(err) => {
msg::write_error_response(out, &PgError::from(&err));
ok = false;
break;
}
}
}
if capture_session {
self.apply_session_updates(let_vars, trailing);
}
if !ok {
break;
}
}
Ok(())
}
async fn run_in_transaction(&mut self, sql: &str, out: &mut BytesMut) -> Result<()> {
let tx = self.transaction.clone().expect("caller checked transaction is open");
let cancel = self.arm_cancel();
let results = self
.ds
.run(QueryRequest::new(sql, &self.session).with_transaction(tx).with_cancel(cancel))
.await;
match results {
Ok(results) => {
for result in results {
match result.result {
Ok(value) => {
if !self.emit_result(value, out).await? {
self.txn_failed = true;
break;
}
}
Err(err) => {
msg::write_error_response(out, &PgError::from(&err));
self.txn_failed = true;
break;
}
}
}
}
Err(err) => {
msg::write_error_response(out, &PgError::from(&err));
self.txn_failed = true;
}
}
Ok(())
}
async fn run_gql(&mut self, sql: &str, out: &mut BytesMut) -> Result<()> {
let cancel = self.arm_cancel();
match self
.ds
.run(QueryRequest::new(QuerySource::gql(sql), &self.session).with_cancel(cancel))
.await
{
Ok(results) => {
for result in results {
match result.result {
Ok(value) => {
if !self.emit_result(value, out).await? {
break;
}
}
Err(err) => {
msg::write_error_response(out, &PgError::from(&err));
break;
}
}
}
}
Err(err) => msg::write_error_response(out, &PgError::from(&err)),
}
Ok(())
}
async fn handle_transaction_control(
&mut self,
op: TxnOp,
out: &mut BytesMut,
) -> Result<(), PgError> {
match op {
TxnOp::Begin => {
if self.transaction.is_some() {
msg::write_notice(out, "25001", "there is already a transaction in progress");
} else {
let tx = self
.ds
.transaction(TransactionType::Write)
.await
.map_err(|err| PgError::internal(err.to_string()))?;
self.transaction = Some(Arc::new(tx));
self.txn_failed = false;
}
msg::write_command_complete(out, "BEGIN");
}
TxnOp::Commit => match self.transaction.take() {
Some(tx) if self.txn_failed => {
self.txn_failed = false;
let _ = tx.cancel().await;
msg::write_command_complete(out, "ROLLBACK");
}
Some(tx) => {
tx.commit().await.map_err(|err| PgError::internal(err.to_string()))?;
msg::write_command_complete(out, "COMMIT");
}
None => {
msg::write_notice(out, "25P01", "there is no transaction in progress");
msg::write_command_complete(out, "COMMIT");
}
},
TxnOp::Rollback => {
match self.transaction.take() {
Some(tx) => {
let _ = tx.cancel().await;
}
None => {
msg::write_notice(out, "25P01", "there is no transaction in progress");
}
}
self.txn_failed = false;
msg::write_command_complete(out, "ROLLBACK");
}
}
Ok(())
}
fn select_dialect(&mut self, dialect: Dialect) -> Result<(), PgError> {
self.check_dialect(dialect)?;
self.dialect = dialect;
Ok(())
}
fn check_current_dialect(&self) -> Result<(), PgError> {
self.check_dialect(self.dialect)
}
fn check_dialect(&self, dialect: Dialect) -> Result<(), PgError> {
let allowed = match dialect {
Dialect::SurrealQl => true,
Dialect::Gql => self.ds.allows_http_route(&RouteTarget::Gql),
};
if !allowed {
let name = dialect_name(dialect);
warn!(
target: LOG,
"Capabilities denied postgres {name} dialect request from {}", self.peer
);
return Err(PgError::insufficient_privilege(format!(
"Forbidden: the {name} dialect is not allowed on this server"
)));
}
Ok(())
}
fn should_intercept_set(&self, set: &SetStatement) -> bool {
self.dialect == Dialect::SurrealQl
|| set.name.is_empty()
|| set.name.eq_ignore_ascii_case("dialect")
|| is_no_op_guc(&set.name)
}
fn handle_set(&mut self, set: SetStatement, out: &mut BytesMut) -> Result<(), PgError> {
let SetStatement {
name,
value,
} = set;
if name.is_empty() {
msg::write_command_complete(out, "SET");
return Ok(());
}
if name.eq_ignore_ascii_case("dialect") {
let dialect = parse_dialect(&value)
.ok_or_else(|| PgError::invalid_text(format!("unrecognized dialect: {value}")))?;
self.select_dialect(dialect)?;
msg::write_command_complete(out, "SET");
return Ok(());
}
if is_no_op_guc(&name) {
msg::write_command_complete(out, "SET");
Ok(())
} else {
Err(PgError::undefined_object(format!(
"unrecognized configuration parameter \"{name}\""
)))
}
}
fn arm_cancel(&mut self) -> CancelHandle {
let cancel = CancelHandle::new();
self.cancel = cancel.clone();
self.registry.insert((self.pid, self.secret), cancel.clone());
cancel
}
fn is_tls(&self) -> bool {
matches!(self.stream, Either::Right(_))
}
fn ready_status(&self) -> u8 {
if self.transaction.is_none() {
b'I'
} else if self.txn_failed {
b'E'
} else {
b'T'
}
}
fn aborted_txn_guard(&self) -> Option<PgError> {
self.txn_failed.then(|| {
PgError::in_failed_transaction(
"current transaction is aborted, commands ignored until end of transaction block",
)
})
}
fn note_extended_error(&mut self, err: &PgError, out: &mut BytesMut) {
msg::write_error_response(out, err);
if self.transaction.is_some() {
self.txn_failed = true;
}
}
fn note_simple_error(&mut self, err: &PgError, out: &mut BytesMut) {
msg::write_error_response(out, err);
if self.transaction.is_some() {
self.txn_failed = true;
}
}
async fn emit_result(&mut self, value: Value, out: &mut BytesMut) -> Result<bool> {
let ResultShape {
columns,
rows,
} = shape_result(value);
if columns.len() > msg::MAX_COLUMNS {
msg::write_error_response(
out,
&PgError::feature_not_supported(format!(
"result has {} columns, exceeding the maximum of {}",
columns.len(),
msg::MAX_COLUMNS
)),
);
return Ok(false);
}
msg::write_row_description(out, &columns, |_| FORMAT_TEXT);
let mut count = 0usize;
for row in rows {
let cells = match encode_row(row, &columns, &[]) {
Ok(cells) => cells,
Err(err) => {
msg::write_error_response(out, &err);
return Ok(false);
}
};
msg::write_data_row(out, &cells);
count += 1;
if out.len() >= FLUSH_THRESHOLD {
self.stream.write_all(out).await?;
out.clear();
}
}
msg::write_command_complete(out, &format!("SELECT {count}"));
Ok(true)
}
fn apply_session_updates(
&mut self,
let_vars: Vec<String>,
mut trailing: Vec<surrealdb_rpc::QueryResult>,
) {
let Some(session_result) = trailing.pop() else {
return;
};
for (name, result) in let_vars.into_iter().zip(trailing) {
if surrealdb_rpc::check_protected_param(&name).is_err() {
continue;
}
if let Ok(value) = result.result {
self.session.variables.insert(name, value);
}
}
if let Ok(Value::Object(obj)) = session_result.result {
let mut map = obj.into_inner();
self.session.ns = match map.remove("ns") {
Some(Value::String(ns)) => Some(ns),
_ => None,
};
self.session.db = match map.remove("db") {
Some(Value::String(db)) => Some(db),
_ => None,
};
}
}
async fn finish_query(&mut self, mut out: BytesMut) -> Result<()> {
let status = self.ready_status();
msg::write_ready_for_query(&mut out, status);
self.stream.write_all(&out).await?;
self.stream.flush().await?;
Ok(())
}
async fn report(&mut self, err: &PgError) -> Result<()> {
let mut buf = BytesMut::new();
msg::write_error_response(&mut buf, err);
self.stream.write_all(&buf).await?;
self.stream.flush().await?;
Ok(())
}
async fn read_startup_packet(&mut self) -> Result<Option<Vec<u8>>> {
let len = tokio::select! {
biased;
_ = self.shutdown.cancelled() => return Ok(None),
len = self.stream.read_i32() => match len {
Ok(len) => len,
Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None),
Err(err) => return Err(err.into()),
},
};
match read_startup_payload(&mut self.stream, len).await? {
Some(payload) => Ok(Some(payload)),
None => {
let _ =
self.report(&PgError::protocol("invalid startup packet length").fatal()).await;
Ok(None)
}
}
}
async fn read_message(&mut self, max: usize) -> Result<Option<(u8, Vec<u8>)>> {
let tag = tokio::select! {
biased;
_ = self.shutdown.cancelled() => return Ok(None),
tag = self.stream.read_u8() => match tag {
Ok(tag) => tag,
Err(err) if err.kind() == std::io::ErrorKind::UnexpectedEof => return Ok(None),
Err(err) => return Err(err.into()),
},
};
let len = self.stream.read_i32().await?;
let len = usize::try_from(len).unwrap_or(0);
if !(4..=max).contains(&len) {
let _ = self.report(&PgError::protocol("invalid message length").fatal()).await;
return Ok(None);
}
let mut payload = vec![0u8; len - 4];
self.stream.read_exact(&mut payload).await?;
Ok(Some((tag, payload)))
}
}
fn rewrite_positional_params(query: &str) -> (String, usize) {
#[derive(PartialEq)]
enum State {
Normal,
Single,
Double,
Backtick,
Angle,
Line,
Block,
}
let chars: Vec<char> = query.chars().collect();
let mut out = String::with_capacity(query.len() + 8);
let mut state = State::Normal;
let mut max_index = 0usize;
let mut i = 0;
while i < chars.len() {
let c = chars[i];
let peek = chars.get(i + 1).copied();
match state {
State::Normal => match c {
'\'' => {
state = State::Single;
out.push(c);
i += 1;
}
'"' => {
state = State::Double;
out.push(c);
i += 1;
}
'`' => {
state = State::Backtick;
out.push(c);
i += 1;
}
'\u{27e8}' => {
state = State::Angle;
out.push(c);
i += 1;
}
'-' if peek == Some('-') => {
state = State::Line;
out.push_str("--");
i += 2;
}
'#' => {
state = State::Line;
out.push('#');
i += 1;
}
'/' if peek == Some('*') => {
state = State::Block;
out.push_str("/*");
i += 2;
}
'$' if peek.is_some_and(|p| p.is_ascii_digit()) => {
out.push_str("$_");
i += 1;
let mut digits = String::new();
while i < chars.len() && chars[i].is_ascii_digit() {
digits.push(chars[i]);
out.push(chars[i]);
i += 1;
}
if let Ok(index) = digits.parse::<usize>() {
max_index = max_index.max(index);
}
}
_ => {
out.push(c);
i += 1;
}
},
State::Single | State::Double => {
out.push(c);
i += 1;
if c == '\\' {
if let Some(next) = chars.get(i) {
out.push(*next);
i += 1;
}
} else if (state == State::Single && c == '\'')
|| (state == State::Double && c == '"')
{
state = State::Normal;
}
}
State::Backtick => {
out.push(c);
i += 1;
if c == '`' {
state = State::Normal;
}
}
State::Angle => {
out.push(c);
i += 1;
if c == '\u{27e9}' {
state = State::Normal;
}
}
State::Line => {
out.push(c);
i += 1;
if c == '\n' {
state = State::Normal;
}
}
State::Block => {
out.push(c);
i += 1;
if c == '*' && chars.get(i) == Some(&'/') {
out.push('/');
i += 1;
state = State::Normal;
}
}
}
}
(out, max_index)
}
fn format_at(formats: &[i16], i: usize) -> i16 {
match formats.len() {
0 => FORMAT_TEXT,
1 => formats[0],
_ => formats.get(i).copied().unwrap_or(FORMAT_TEXT),
}
}
fn row_within_wire_limit(cells: &[Option<Vec<u8>>]) -> bool {
wire_row_fits(cells.iter().map(|c| c.as_ref().map_or(0, Vec::len)))
}
fn wire_row_fits(cell_lengths: impl Iterator<Item = usize>) -> bool {
let mut total: u64 = 6;
for len in cell_lengths {
total += 4 + len as u64;
if total > i32::MAX as u64 {
return false;
}
}
true
}
fn encode_cell(value: Value, ty: PgType, format: i16) -> Result<Vec<u8>, PgError> {
if format == FORMAT_BINARY {
encode::encode_binary(value, ty)
} else {
encode::encode_text(value, ty)
}
}
fn encode_row(
row: Vec<Option<Value>>,
columns: &[PgColumn],
formats: &[i16],
) -> Result<Vec<Option<Vec<u8>>>, PgError> {
let mut cells = Vec::with_capacity(row.len());
for (i, (cell, column)) in row.into_iter().zip(columns).enumerate() {
match cell {
Some(value) => cells.push(Some(encode_cell(value, column.ty, format_at(formats, i))?)),
None => cells.push(None),
}
}
if !row_within_wire_limit(&cells) {
return Err(PgError::feature_not_supported(
"result row exceeds the maximum Postgres wire size",
));
}
Ok(cells)
}
fn parse_database_param(user: &str, database: Option<&str>) -> (Option<String>, Option<String>) {
let non_empty = |s: &str| {
if s.is_empty() {
None
} else {
Some(s.to_string())
}
};
match database {
None => (None, None),
Some(database) => match database.split_once('/') {
Some((ns, db)) => (non_empty(ns), non_empty(db)),
None if database == user => (None, None),
None => (non_empty(database), None),
},
}
}
fn auth_failed(user: &str) -> PgError {
PgError::invalid_password(format!("password authentication failed for user \"{user}\"")).fatal()
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
enum TxnOp {
Begin,
Commit,
Rollback,
}
fn transaction_control(trimmed: &str) -> Option<TxnOp> {
let stripped = strip_sql_comments(trimmed);
let stripped = stripped.trim().trim_end_matches(';').trim();
if stripped.contains(';') {
return None;
}
let mut tokens = stripped.split_whitespace().map(str::to_ascii_uppercase);
let (op, start) = match tokens.next()?.as_str() {
"BEGIN" => (TxnOp::Begin, false),
"START" => (TxnOp::Begin, true), "COMMIT" | "END" => (TxnOp::Commit, false),
"ROLLBACK" | "ABORT" => (TxnOp::Rollback, false),
_ => return None,
};
if start && tokens.next().as_deref() != Some("TRANSACTION") {
return None;
}
const MODIFIERS: &[&str] = &[
"TRANSACTION",
"WORK",
"ISOLATION",
"LEVEL",
"SERIALIZABLE",
"REPEATABLE",
"READ",
"WRITE",
"ONLY",
"COMMITTED",
"UNCOMMITTED",
"DEFERRABLE",
"NOT",
"AND",
"NO",
"CHAIN",
];
if tokens.all(|t| MODIFIERS.contains(&t.as_str())) {
Some(op)
} else {
None
}
}
fn strip_sql_comments(sql: &str) -> String {
let bytes = sql.as_bytes();
let mut out = String::with_capacity(sql.len());
let mut i = 0;
while i < bytes.len() {
match bytes[i] {
b'-' if bytes.get(i + 1) == Some(&b'-') => {
while i < bytes.len() && bytes[i] != b'\n' {
i += 1;
}
}
b'#' => {
while i < bytes.len() && bytes[i] != b'\n' {
i += 1;
}
}
b'/' if bytes.get(i + 1) == Some(&b'*') => {
i += 2;
while i < bytes.len() && !(bytes[i] == b'*' && bytes.get(i + 1) == Some(&b'/')) {
i += 1;
}
i += 2;
out.push(' ');
}
_ => {
let ch = sql[i..].chars().next().expect("valid char boundary");
out.push(ch);
i += ch.len_utf8();
}
}
}
out
}
#[derive(Clone)]
struct SetStatement {
name: String,
value: String,
}
const NO_OP_GUCS: &[&str] = &[
"extra_float_digits",
"application_name",
"client_encoding",
"datestyle",
"timezone",
"statement_timeout",
"search_path",
"standard_conforming_strings",
"client_min_messages",
"bytea_output",
"intervalstyle",
];
fn is_no_op_guc(name: &str) -> bool {
NO_OP_GUCS.iter().any(|g| name.eq_ignore_ascii_case(g))
}
fn parse_set(trimmed: &str) -> Option<SetStatement> {
let rest = trimmed.trim_end_matches(';').trim();
let mut words = rest.splitn(2, char::is_whitespace);
if !words.next()?.eq_ignore_ascii_case("SET") {
return None;
}
let assignment = words.next().unwrap_or("").trim();
let split = assignment.find('=').map(|i| (assignment[..i].trim(), assignment[i + 1..].trim()));
let split = split.or_else(|| {
assignment
.to_ascii_uppercase()
.find(" TO ")
.map(|i| (assignment[..i].trim(), assignment[i + 4..].trim()))
});
Some(match split {
Some((name, value)) => SetStatement {
name: name.to_string(),
value: value.to_string(),
},
None => SetStatement {
name: String::new(),
value: String::new(),
},
})
}
fn dialect_from_options(options: &str) -> Option<Dialect> {
let mut tokens = options.split_whitespace().peekable();
while let Some(token) = tokens.next() {
let setting = if token == "-c" {
tokens.next()?
} else {
token.strip_prefix("-c").unwrap_or(token)
};
if let Some(value) = setting.strip_prefix("dialect=") {
return parse_dialect(value);
}
}
None
}
fn parse_dialect(value: &str) -> Option<Dialect> {
match value.trim().trim_matches('\'').to_ascii_lowercase().as_str() {
"surrealql" | "surql" | "sql" => Some(Dialect::SurrealQl),
"gql" => Some(Dialect::Gql),
_ => None,
}
}
fn dialect_name(dialect: Dialect) -> &'static str {
match dialect {
Dialect::SurrealQl => "surrealql",
Dialect::Gql => "gql",
}
}
#[cfg(test)]
mod tests {
use super::{
Dialect, TxnOp, dialect_from_options, dialect_name, parse_database_param, parse_dialect,
parse_set, rewrite_positional_params, row_within_wire_limit, transaction_control,
wire_row_fits,
};
fn rewrite(query: &str) -> String {
rewrite_positional_params(query).0
}
#[test]
fn recognises_transaction_control() {
assert_eq!(transaction_control("BEGIN"), Some(TxnOp::Begin));
assert_eq!(transaction_control("begin;"), Some(TxnOp::Begin));
assert_eq!(transaction_control("START TRANSACTION"), Some(TxnOp::Begin));
assert_eq!(transaction_control("COMMIT"), Some(TxnOp::Commit));
assert_eq!(transaction_control("End Transaction ;"), Some(TxnOp::Commit));
assert_eq!(transaction_control("ROLLBACK"), Some(TxnOp::Rollback));
assert_eq!(transaction_control("ABORT"), Some(TxnOp::Rollback));
assert_eq!(transaction_control("BEGIN; RETURN 1; COMMIT"), None);
assert_eq!(transaction_control("RETURN 1"), None);
}
#[test]
fn recognises_transaction_preambles() {
assert_eq!(transaction_control("BEGIN ISOLATION LEVEL SERIALIZABLE"), Some(TxnOp::Begin));
assert_eq!(transaction_control("START TRANSACTION READ WRITE"), Some(TxnOp::Begin));
assert_eq!(transaction_control("BEGIN /* jdbc */"), Some(TxnOp::Begin));
assert_eq!(transaction_control("ROLLBACK AND NO CHAIN"), Some(TxnOp::Rollback));
assert_eq!(transaction_control("START"), None);
assert_eq!(transaction_control("BEGIN something weird"), None);
assert_eq!(transaction_control("COMMIT; SELECT 1"), None);
}
#[test]
fn row_wire_limit() {
assert!(row_within_wire_limit(&[Some(vec![1, 2, 3]), None]));
assert!(wire_row_fits(std::iter::once(i32::MAX as usize - 16)));
assert!(!wire_row_fits(std::iter::once(i32::MAX as usize)));
assert!(!wire_row_fits(std::iter::repeat_n(1_000_000_000, 3)));
}
#[test]
fn parses_set_statements() {
let set = parse_set("SET dialect = 'gql'").unwrap();
assert_eq!(set.name, "dialect");
assert_eq!(set.value, "'gql'");
let set = parse_set("SET extra_float_digits TO 3").unwrap();
assert_eq!(set.name, "extra_float_digits");
assert_eq!(set.value, "3");
assert_eq!(parse_set("SET TIME ZONE 'UTC'").unwrap().name, "");
assert!(parse_set("SELECT 1").is_none());
}
#[test]
fn parses_dialects() {
assert_eq!(parse_dialect("gql"), Some(Dialect::Gql));
assert_eq!(parse_dialect("'GQL'"), Some(Dialect::Gql));
assert_eq!(parse_dialect("surrealql"), Some(Dialect::SurrealQl));
assert_eq!(parse_dialect("nope"), None);
}
#[test]
fn dialect_names_round_trip() {
for dialect in [Dialect::SurrealQl, Dialect::Gql] {
assert_eq!(parse_dialect(dialect_name(dialect)), Some(dialect));
}
}
#[test]
fn parses_dialect_from_options() {
assert_eq!(dialect_from_options("-c dialect=gql"), Some(Dialect::Gql));
assert_eq!(dialect_from_options("-cdialect=gql"), Some(Dialect::Gql));
assert_eq!(dialect_from_options("dialect=surrealql"), Some(Dialect::SurrealQl));
assert_eq!(dialect_from_options("-c statement_timeout=5000"), None);
}
#[test]
fn rewrites_positional_params() {
assert_eq!(rewrite("RETURN $1 + $2"), "RETURN $_1 + $_2");
assert_eq!(rewrite("SELECT * FROM t WHERE id = $10"), "SELECT * FROM t WHERE id = $_10");
}
#[test]
fn reports_highest_positional_index() {
assert_eq!(rewrite_positional_params("RETURN $1 + $3").1, 3);
assert_eq!(rewrite_positional_params("RETURN 'no $9 here'").1, 0);
assert_eq!(rewrite_positional_params("RETURN $name").1, 0);
}
#[test]
fn leaves_named_params_untouched() {
assert_eq!(rewrite("RETURN $name"), "RETURN $name");
assert_eq!(rewrite("RETURN $_1"), "RETURN $_1");
}
#[test]
fn does_not_rewrite_inside_strings() {
assert_eq!(rewrite("RETURN 'price $1'"), "RETURN 'price $1'");
assert_eq!(rewrite(r#"RETURN "cost $2""#), r#"RETURN "cost $2""#);
assert_eq!(rewrite(r#"RETURN 'a\'b $1' + $2"#), r#"RETURN 'a\'b $1' + $_2"#);
}
#[test]
fn does_not_rewrite_inside_angle_identifiers() {
assert_eq!(
rewrite("SELECT * FROM ⟨tbl$1⟩ WHERE x = $2"),
"SELECT * FROM ⟨tbl$1⟩ WHERE x = $_2"
);
assert_eq!(rewrite("RETURN $⟨name⟩"), "RETURN $⟨name⟩");
}
#[test]
fn does_not_rewrite_inside_comments() {
assert_eq!(rewrite("RETURN 1 -- $1\n+ $2"), "RETURN 1 -- $1\n+ $_2");
assert_eq!(rewrite("RETURN /* $1 */ $2"), "RETURN /* $1 */ $_2");
}
#[test]
fn database_param_parsing() {
let s = |v: &str| Some(v.to_string());
assert_eq!(parse_database_param("root", None), (None, None));
assert_eq!(parse_database_param("root", Some("root")), (None, None));
assert_eq!(parse_database_param("root", Some("")), (None, None));
assert_eq!(parse_database_param("root", Some("ns")), (s("ns"), None));
assert_eq!(parse_database_param("root", Some("ns/")), (s("ns"), None));
assert_eq!(parse_database_param("root", Some("ns/db")), (s("ns"), s("db")));
assert_eq!(parse_database_param("root", Some("/db")), (None, s("db")));
assert_eq!(parse_database_param("root", Some("root/db")), (s("root"), s("db")));
}
}