mod protocol;
mod types;
use std::collections::{BTreeMap, HashMap};
use std::future::Future;
use std::sync::Arc;
use corium_core::EntityId;
use corium_db::Db;
use corium_query::edn::Edn;
use corium_sql::{MutationKind, SqlColumn, SqlError, SqlSession, SqlType};
use thiserror::Error;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio::net::TcpListener;
use protocol::{BackendWriter, ErrorFields, FieldDescription, Frontend, FrontendReader};
#[derive(Debug, Error)]
pub enum CatalogError {
#[error("database {0:?} is not available")]
NotFound(String),
#[error("{0}")]
Unavailable(String),
#[error("database {0:?} is read-only through this catalog")]
ReadOnly(String),
#[error("{0}")]
Conflict(String),
#[error("{0}")]
Rejected(String),
#[error("{0}")]
Denied(String),
#[error("{0}")]
Unsupported(String),
}
pub struct CatalogTxResult {
pub db_after: Db,
pub tempids: BTreeMap<String, EntityId>,
}
#[async_trait::async_trait]
pub trait DbCatalog: Send + Sync + 'static {
async fn list(&self) -> Result<Vec<String>, CatalogError>;
async fn db(&self, name: &str) -> Result<Db, CatalogError>;
async fn transact(
&self,
name: &str,
_expected_basis_t: u64,
_forms: Vec<Edn>,
) -> Result<CatalogTxResult, CatalogError> {
Err(CatalogError::ReadOnly(name.to_owned()))
}
}
#[derive(Clone, Debug)]
pub struct PgWireConfig {
pub password: Option<String>,
pub server_version: String,
}
impl Default for PgWireConfig {
fn default() -> Self {
Self {
password: None,
server_version: concat!("16.0 (corium ", env!("CARGO_PKG_VERSION"), ")").to_owned(),
}
}
}
pub async fn serve<C, F>(
listener: TcpListener,
catalog: Arc<C>,
config: PgWireConfig,
shutdown: F,
) -> std::io::Result<()>
where
C: DbCatalog,
F: Future<Output = ()>,
{
let config = Arc::new(config);
tokio::pin!(shutdown);
loop {
tokio::select! {
() = &mut shutdown => return Ok(()),
accepted = listener.accept() => {
let (stream, peer) = accepted?;
let catalog = Arc::clone(&catalog);
let config = Arc::clone(&config);
tokio::spawn(async move {
let (read, write) = stream.into_split();
let mut session = ConnectionSession::new(
FrontendReader::new(read),
BackendWriter::new(write),
catalog,
config,
);
if let Err(error) = session.run().await {
tracing::debug!(%peer, %error, "pgwire connection closed");
}
});
}
}
}
}
struct Portal {
sql: String,
database: Option<String>,
params: Vec<corium_sql::SqlValue>,
result_formats: Vec<i16>,
}
struct PreparedStatement {
sql: String,
parameter_types: Vec<i32>,
}
enum Statement {
Control(&'static str),
Begin,
Commit,
Rollback,
Use(String),
ShowDatabases,
Query,
Mutation,
}
enum Dispatch {
NoDatabase,
Catalog(CatalogError),
Sql(SqlError),
}
struct ConnectionSession<R, W, C> {
reader: FrontendReader<R>,
writer: BackendWriter<W>,
catalog: Arc<C>,
config: Arc<PgWireConfig>,
current_db: Option<String>,
statements: HashMap<String, PreparedStatement>,
portals: HashMap<String, Portal>,
failed: bool,
in_transaction: bool,
transaction_failed: bool,
transaction_db: Option<(String, Db)>,
}
impl<R, W, C> ConnectionSession<R, W, C>
where
R: AsyncRead + Unpin,
W: AsyncWrite + Unpin,
C: DbCatalog,
{
fn new(
reader: FrontendReader<R>,
writer: BackendWriter<W>,
catalog: Arc<C>,
config: Arc<PgWireConfig>,
) -> Self {
Self {
reader,
writer,
catalog,
config,
current_db: None,
statements: HashMap::new(),
portals: HashMap::new(),
failed: false,
in_transaction: false,
transaction_failed: false,
transaction_db: None,
}
}
async fn run(&mut self) -> std::io::Result<()> {
let startup = self.reader.read_startup(&mut self.writer).await?;
if !self.authenticate().await? {
return Ok(());
}
self.current_db = startup
.get("database")
.filter(|value| !value.is_empty())
.map(str::to_owned);
self.send_ready_banner(startup.get("application_name").unwrap_or(""))
.await?;
while let Some(message) = self.reader.read_message().await? {
match message {
Frontend::Query(sql) => {
if !self.failed {
self.simple_query(&sql).await?;
self.writer.flush().await?;
}
}
Frontend::Parse {
name,
query,
parameter_types,
} => self.handle_parse(name, query, parameter_types),
Frontend::Bind {
portal,
statement,
parameter_formats,
parameters,
result_formats,
} => self.handle_bind(
&portal,
&statement,
¶meter_formats,
¶meters,
&result_formats,
),
Frontend::Describe { kind, name } => self.handle_describe(kind, &name).await?,
Frontend::Execute { portal } => self.handle_execute(&portal).await?,
Frontend::Close { kind, name } => {
if !self.failed {
if kind == b'S' {
self.statements.remove(&name);
} else {
self.portals.remove(&name);
}
self.writer.close_complete();
}
}
Frontend::Sync => {
self.failed = false;
self.writer.ready_for_query(self.ready_status());
self.writer.flush().await?;
}
Frontend::Flush => self.writer.flush().await?,
Frontend::Password(_) => {}
Frontend::Terminate => break,
}
}
Ok(())
}
async fn authenticate(&mut self) -> std::io::Result<bool> {
let Some(expected) = self.config.password.clone() else {
self.writer.authentication_ok();
return Ok(true);
};
self.writer.authentication_cleartext_password();
self.writer.flush().await?;
match self.reader.read_message().await? {
Some(Frontend::Password(supplied)) if supplied == expected => {
self.writer.authentication_ok();
Ok(true)
}
_ => {
self.writer.error_response(&ErrorFields {
code: "28P01",
message: "password authentication failed",
});
self.writer.flush().await?;
Ok(false)
}
}
}
async fn send_ready_banner(&mut self, application_name: &str) -> std::io::Result<()> {
self.writer
.parameter_status("server_version", &self.config.server_version);
self.writer.parameter_status("server_encoding", "UTF8");
self.writer.parameter_status("client_encoding", "UTF8");
self.writer.parameter_status("DateStyle", "ISO, MDY");
self.writer.parameter_status("TimeZone", "UTC");
self.writer.parameter_status("integer_datetimes", "on");
self.writer
.parameter_status("standard_conforming_strings", "on");
self.writer
.parameter_status("application_name", application_name);
self.writer.backend_key_data(0, 0);
self.writer.ready_for_query(b'I');
self.writer.flush().await
}
async fn simple_query(&mut self, sql: &str) -> std::io::Result<()> {
let statements = split_statements(sql);
if statements.is_empty() {
self.writer.empty_query_response();
self.writer.ready_for_query(self.ready_status());
return Ok(());
}
for statement in statements {
if !self.run_simple_statement(&statement).await? {
break;
}
}
self.writer.ready_for_query(self.ready_status());
Ok(())
}
async fn run_simple_statement(&mut self, sql: &str) -> std::io::Result<bool> {
let statement = classify(sql);
if self.transaction_failed && !matches!(&statement, Statement::Rollback) {
self.report_transaction_aborted();
return Ok(false);
}
match statement {
Statement::Control(tag) => {
self.writer.command_complete(tag);
Ok(true)
}
Statement::Begin => {
self.begin_transaction();
self.writer.command_complete("BEGIN");
Ok(true)
}
Statement::Commit => {
self.end_transaction();
self.writer.command_complete("COMMIT");
Ok(true)
}
Statement::Rollback => {
self.end_transaction();
self.writer.command_complete("ROLLBACK");
Ok(true)
}
Statement::Use(name) => match self.use_database(&name).await {
Ok(()) => {
self.writer.command_complete("USE");
Ok(true)
}
Err(error) => {
self.report_dispatch(&error);
Ok(false)
}
},
Statement::ShowDatabases => match self.show_databases(true, &[]).await {
Ok(()) => Ok(true),
Err(error) => {
self.report_dispatch(&error);
Ok(false)
}
},
Statement::Query => {
let db = match self.snapshot(None).await {
Ok(db) => db,
Err(error) => {
self.report_dispatch(&error);
return Ok(false);
}
};
match self.run_statement(&db, sql, &[], true, &[]).await {
Ok(rows) => {
self.writer.command_complete(&command_tag(sql, rows));
Ok(true)
}
Err(error) => {
self.report_dispatch(&Dispatch::Sql(error));
Ok(false)
}
}
}
Statement::Mutation => {
if self.in_transaction {
self.report_dispatch(&explicit_transaction_write_error());
return Ok(false);
}
let Some(database) = self.current_db.clone() else {
self.report_dispatch(&Dispatch::NoDatabase);
return Ok(false);
};
let db = match self.snapshot(Some(&database)).await {
Ok(db) => db,
Err(error) => {
self.report_dispatch(&error);
return Ok(false);
}
};
match self.run_mutation(&database, &db, sql, &[], true, &[]).await {
Ok((kind, rows)) => {
self.writer
.command_complete(&mutation_command_tag(kind, rows));
Ok(true)
}
Err(error) => {
self.report_dispatch(&error);
Ok(false)
}
}
}
}
}
fn handle_parse(&mut self, name: String, query: String, mut parameter_types: Vec<i32>) {
if self.failed {
return;
}
let inferred_count = placeholder_count(&query);
if parameter_types.len() > inferred_count {
self.fail_extended("08P01", "too many parameter types in Parse");
return;
}
parameter_types.resize(inferred_count, 0);
self.statements.insert(
name,
PreparedStatement {
sql: query,
parameter_types,
},
);
self.writer.parse_complete();
}
fn handle_bind(
&mut self,
portal: &str,
statement: &str,
parameter_formats: &[i16],
parameters: &[Option<Vec<u8>>],
result_formats: &[i16],
) {
if self.failed {
return;
}
if result_formats.iter().any(|format| !matches!(format, 0 | 1)) {
self.fail_extended("08P01", "result format code must be zero or one");
return;
}
let Some(prepared) = self.statements.get(statement) else {
self.fail_extended("26000", "prepared statement does not exist");
return;
};
if parameters.len() != prepared.parameter_types.len() {
self.fail_extended("08P01", "bound parameter count does not match Parse");
return;
}
let formats = match expand_formats(parameter_formats, parameters.len()) {
Ok(formats) => formats,
Err(message) => {
self.fail_extended("08P01", message);
return;
}
};
let params = prepared
.parameter_types
.iter()
.zip(formats)
.zip(parameters)
.map(|((oid, format), value)| types::decode_parameter(*oid, format, value.as_deref()))
.collect::<Result<Vec<_>, _>>();
let params = match params {
Ok(params) => params,
Err(message) => {
self.fail_extended("22P02", &message);
return;
}
};
self.portals.insert(
portal.to_owned(),
Portal {
sql: prepared.sql.clone(),
database: self.current_db.clone(),
params,
result_formats: result_formats.to_vec(),
},
);
self.writer.bind_complete();
}
async fn handle_describe(&mut self, kind: u8, name: &str) -> std::io::Result<()> {
if self.failed {
return Ok(());
}
let (sql, database, params, result_formats) = if kind == b'S' {
let Some((sql, parameter_types)) = self
.statements
.get(name)
.map(|statement| (statement.sql.clone(), statement.parameter_types.clone()))
else {
self.fail_extended("26000", "prepared statement does not exist");
return Ok(());
};
self.writer.parameter_description(¶meter_types);
let params = parameter_types
.into_iter()
.map(types::describe_parameter)
.collect::<Result<Vec<_>, _>>();
let params = match params {
Ok(params) => params,
Err(message) => {
self.fail_extended("0A000", &message);
return Ok(());
}
};
(sql, self.current_db.clone(), params, Vec::new())
} else {
let Some(portal) = self.portals.get(name) else {
self.fail_extended("34000", "portal does not exist");
return Ok(());
};
(
portal.sql.clone(),
portal.database.clone(),
portal.params.clone(),
portal.result_formats.clone(),
)
};
match classify(&sql) {
Statement::ShowDatabases => {
self.write_row_description(&[database_field()], &result_formats);
}
Statement::Query if !sql.trim().is_empty() => {
let db = match self.snapshot(database.as_deref()).await {
Ok(db) => db,
Err(error) => {
self.fail_dispatch(&error);
return Ok(());
}
};
match self.describe_columns(&db, &sql, ¶ms).await {
Ok(fields) => {
self.write_row_description(&fields, &result_formats);
}
Err(error) => self.fail_dispatch(&Dispatch::Sql(error)),
}
}
Statement::Mutation => {
let db = match self.snapshot(database.as_deref()).await {
Ok(db) => db,
Err(error) => {
self.fail_dispatch(&error);
return Ok(());
}
};
match SqlSession::new(&db) {
Ok(session) => match session.mutation_columns(&sql, ¶ms).await {
Ok(Some(columns)) if columns.is_empty() => self.writer.no_data(),
Ok(Some(columns)) => {
self.write_row_description(
&columns.iter().map(field_of).collect::<Vec<_>>(),
&result_formats,
);
}
Ok(None) => self.writer.no_data(),
Err(error) => self.fail_dispatch(&Dispatch::Sql(error)),
},
Err(error) => self.fail_dispatch(&Dispatch::Sql(error)),
}
}
_ => self.writer.no_data(),
}
Ok(())
}
async fn handle_execute(&mut self, portal: &str) -> std::io::Result<()> {
if self.failed {
return Ok(());
}
let Some((sql, database, params, result_formats)) =
self.portals.get(portal).map(|portal| {
(
portal.sql.clone(),
portal.database.clone(),
portal.params.clone(),
portal.result_formats.clone(),
)
})
else {
self.fail_extended("34000", "portal does not exist");
return Ok(());
};
if sql.trim().is_empty() {
self.writer.empty_query_response();
return Ok(());
}
let statement = classify(&sql);
if self.transaction_failed && !matches!(&statement, Statement::Rollback) {
self.fail_extended(
"25P02",
"current transaction is aborted; commands ignored until end of transaction block",
);
return Ok(());
}
match statement {
Statement::Control(tag) => self.writer.command_complete(tag),
Statement::Begin => {
self.begin_transaction();
self.writer.command_complete("BEGIN");
}
Statement::Commit => {
self.end_transaction();
self.writer.command_complete("COMMIT");
}
Statement::Rollback => {
self.end_transaction();
self.writer.command_complete("ROLLBACK");
}
Statement::Use(name) => match self.use_database(&name).await {
Ok(()) => self.writer.command_complete("USE"),
Err(error) => self.fail_dispatch(&error),
},
Statement::ShowDatabases => {
if let Err(error) = self.show_databases(false, &result_formats).await {
self.fail_dispatch(&error);
}
}
Statement::Query => {
let db = match self.snapshot(database.as_deref()).await {
Ok(db) => db,
Err(error) => {
self.fail_dispatch(&error);
return Ok(());
}
};
match self
.run_statement(&db, &sql, ¶ms, false, &result_formats)
.await
{
Ok(rows) => self.writer.command_complete(&command_tag(&sql, rows)),
Err(error) => self.fail_dispatch(&Dispatch::Sql(error)),
}
}
Statement::Mutation => {
if self.in_transaction {
self.fail_dispatch(&explicit_transaction_write_error());
return Ok(());
}
let Some(database) = database.or_else(|| self.current_db.clone()) else {
self.fail_dispatch(&Dispatch::NoDatabase);
return Ok(());
};
let db = match self.snapshot(Some(&database)).await {
Ok(db) => db,
Err(error) => {
self.fail_dispatch(&error);
return Ok(());
}
};
match self
.run_mutation(&database, &db, &sql, ¶ms, false, &result_formats)
.await
{
Ok((kind, rows)) => self
.writer
.command_complete(&mutation_command_tag(kind, rows)),
Err(error) => self.fail_dispatch(&error),
}
}
}
Ok(())
}
async fn use_database(&mut self, name: &str) -> Result<(), Dispatch> {
self.snapshot(Some(name)).await?;
self.current_db = Some(name.to_owned());
Ok(())
}
async fn show_databases(
&mut self,
with_row_description: bool,
result_formats: &[i16],
) -> Result<(), Dispatch> {
let names = self.catalog.list().await.map_err(Dispatch::Catalog)?;
if with_row_description {
self.write_row_description(&[database_field()], result_formats);
}
for name in &names {
self.writer.data_row(&[Some(name.clone().into_bytes())]);
}
self.writer.command_complete("SHOW");
Ok(())
}
async fn snapshot(&mut self, database: Option<&str>) -> Result<Db, Dispatch> {
let name = database
.map(str::to_owned)
.or_else(|| self.current_db.clone())
.ok_or(Dispatch::NoDatabase)?;
if self.in_transaction {
if let Some((pinned_name, pinned)) = &self.transaction_db {
if pinned_name != &name {
return Err(Dispatch::Sql(SqlError::Mutation(
"cannot switch databases inside an explicit transaction".into(),
)));
}
return Ok(pinned.clone());
}
let db = self.catalog.db(&name).await.map_err(Dispatch::Catalog)?;
self.transaction_db = Some((name, db.clone()));
Ok(db)
} else {
self.catalog.db(&name).await.map_err(Dispatch::Catalog)
}
}
fn ready_status(&self) -> u8 {
if self.transaction_failed {
b'E'
} else if self.in_transaction {
b'T'
} else {
b'I'
}
}
async fn describe_columns(
&self,
db: &Db,
sql: &str,
params: &[corium_sql::SqlValue],
) -> Result<Vec<FieldDescription>, SqlError> {
let session = SqlSession::new(db)?;
let query = session.query_params(sql, params).await?;
Ok(query.columns().iter().map(field_of).collect())
}
async fn run_statement(
&mut self,
db: &Db,
sql: &str,
params: &[corium_sql::SqlValue],
with_row_description: bool,
result_formats: &[i16],
) -> Result<usize, SqlError> {
let session = SqlSession::new(db)?;
let mut query = session.query_params(sql, params).await?;
let columns = query.columns().to_vec();
let formats = expand_result_formats(result_formats, columns.len())
.map_err(|message| SqlError::Mutation(message.into()))?;
if with_row_description {
let fields = columns.iter().map(field_of).collect::<Vec<_>>();
self.writer.row_description_with_formats(&fields, &formats);
}
let mut count = 0usize;
while let Some(row) = query.next_row().await? {
let values = row
.iter()
.zip(&columns)
.zip(&formats)
.map(|((value, column), format)| {
types::encode_result(value, &column.data_type, *format)
.map_err(SqlError::Mutation)
})
.collect::<Result<Vec<_>, _>>()?;
self.writer.data_row(&values);
count += 1;
if count.is_multiple_of(1024) {
self.writer
.flush()
.await
.map_err(|error| SqlError::Schema(error.to_string()))?;
}
}
Ok(count)
}
async fn run_mutation(
&mut self,
database: &str,
db: &Db,
sql: &str,
params: &[corium_sql::SqlValue],
with_row_description: bool,
result_formats: &[i16],
) -> Result<(MutationKind, usize), Dispatch> {
let session = SqlSession::new(db).map_err(Dispatch::Sql)?;
let mutation = session
.mutation_params(sql, params)
.await
.map_err(Dispatch::Sql)?
.ok_or_else(|| Dispatch::Sql(SqlError::Mutation("expected a mutation".into())))?;
let (db_after, tempids) = if mutation.is_empty() {
(db.clone(), BTreeMap::new())
} else {
let result = self
.catalog
.transact(
database,
mutation.expected_basis_t(),
mutation.forms().to_vec(),
)
.await
.map_err(Dispatch::Catalog)?;
(result.db_after, result.tempids)
};
let returned = mutation
.finish(&db_after, &tempids)
.await
.map_err(Dispatch::Sql)?;
let formats = expand_result_formats(result_formats, returned.columns.len())
.map_err(|message| Dispatch::Sql(SqlError::Mutation(message.into())))?;
if with_row_description && !returned.columns.is_empty() {
self.writer.row_description_with_formats(
&returned.columns.iter().map(field_of).collect::<Vec<_>>(),
&formats,
);
}
for row in returned.rows {
let values = row
.iter()
.zip(&returned.columns)
.zip(&formats)
.map(|((value, column), format)| {
types::encode_result(value, &column.data_type, *format)
.map_err(|message| Dispatch::Sql(SqlError::Mutation(message)))
})
.collect::<Result<Vec<_>, _>>()?;
self.writer.data_row(&values);
}
Ok((mutation.kind(), mutation.affected()))
}
fn report_dispatch(&mut self, error: &Dispatch) {
let (code, message) = dispatch_error_fields(error);
self.writer.error_response(&ErrorFields {
code,
message: &message,
});
if self.in_transaction {
self.transaction_failed = true;
}
}
fn fail_dispatch(&mut self, error: &Dispatch) {
self.report_dispatch(error);
self.failed = true;
}
fn fail_extended(&mut self, code: &str, message: &str) {
self.writer.error_response(&ErrorFields { code, message });
self.failed = true;
if self.in_transaction {
self.transaction_failed = true;
}
}
fn report_transaction_aborted(&mut self) {
self.writer.error_response(&ErrorFields {
code: "25P02",
message: "current transaction is aborted; commands ignored until end of transaction block",
});
}
fn begin_transaction(&mut self) {
self.in_transaction = true;
self.transaction_failed = false;
self.transaction_db = None;
}
fn end_transaction(&mut self) {
self.in_transaction = false;
self.transaction_failed = false;
self.transaction_db = None;
}
fn write_row_description(&mut self, fields: &[FieldDescription], requested: &[i16]) {
match expand_result_formats(requested, fields.len()) {
Ok(formats) => self.writer.row_description_with_formats(fields, &formats),
Err(message) => self.fail_extended("08P01", message),
}
}
}
fn database_field() -> FieldDescription {
let type_oid = types::type_oid(&SqlType::Text);
FieldDescription {
name: "database".to_owned(),
type_oid,
type_len: types::type_len(type_oid),
}
}
fn field_of(column: &SqlColumn) -> FieldDescription {
let type_oid = types::type_oid(&column.data_type);
FieldDescription {
name: column.name.clone(),
type_oid,
type_len: types::type_len(type_oid),
}
}
fn dispatch_error_fields(error: &Dispatch) -> (&'static str, String) {
match error {
Dispatch::NoDatabase => (
"3D000",
"no database selected; run \"USE <database>\" first".to_owned(),
),
Dispatch::Catalog(error @ CatalogError::NotFound(_)) => ("3D000", error.to_string()),
Dispatch::Catalog(error @ CatalogError::Unavailable(_)) => ("08006", error.to_string()),
Dispatch::Catalog(error @ CatalogError::ReadOnly(_)) => ("25006", error.to_string()),
Dispatch::Catalog(error @ CatalogError::Conflict(_)) => ("40001", error.to_string()),
Dispatch::Catalog(error @ CatalogError::Rejected(_)) => ("23000", error.to_string()),
Dispatch::Catalog(error @ CatalogError::Denied(_)) => ("42501", error.to_string()),
Dispatch::Catalog(error @ CatalogError::Unsupported(_)) => ("0A000", error.to_string()),
Dispatch::Sql(error) => (sqlstate_for(error), error.to_string()),
}
}
fn explicit_transaction_write_error() -> Dispatch {
Dispatch::Sql(SqlError::Mutation(
"writes inside explicit transaction blocks are not supported; use autocommit".to_owned(),
))
}
fn sqlstate_for(error: &SqlError) -> &'static str {
match error {
SqlError::Schema(_) => "42P01",
SqlError::Mutation(_) => "0A000",
SqlError::Parser(_) | SqlError::DataFusion(_) | SqlError::Arrow(_) => "42601",
}
}
fn command_tag(sql: &str, rows: usize) -> String {
if first_keyword(sql).eq_ignore_ascii_case("explain") {
"EXPLAIN".to_owned()
} else {
format!("SELECT {rows}")
}
}
fn mutation_command_tag(kind: MutationKind, rows: usize) -> String {
match kind {
MutationKind::Insert => format!("INSERT 0 {rows}"),
MutationKind::Update => format!("UPDATE {rows}"),
MutationKind::Delete => format!("DELETE {rows}"),
}
}
fn expand_formats(formats: &[i16], count: usize) -> Result<Vec<i16>, &'static str> {
match formats {
[] => Ok(vec![0; count]),
[format] => Ok(vec![*format; count]),
formats if formats.len() == count => Ok(formats.to_vec()),
_ => Err("parameter format count must be zero, one, or the parameter count"),
}
}
fn expand_result_formats(formats: &[i16], count: usize) -> Result<Vec<i16>, &'static str> {
match formats {
[] => Ok(vec![0; count]),
[format] => Ok(vec![*format; count]),
formats if formats.len() == count => Ok(formats.to_vec()),
_ => Err("result format count must be zero, one, or the result column count"),
}
}
fn placeholder_count(sql: &str) -> usize {
#[derive(Clone, Copy)]
enum State {
Normal,
Single { backslash_escapes: bool },
Double,
LineComment,
BlockComment(usize),
}
let bytes = sql.as_bytes();
let mut state = State::Normal;
let mut count = 0usize;
let mut index = 0usize;
while index < bytes.len() {
match state {
State::Normal if bytes[index..].starts_with(b"--") => {
state = State::LineComment;
index += 1;
}
State::Normal if bytes[index..].starts_with(b"/*") => {
state = State::BlockComment(1);
index += 1;
}
State::Normal if bytes[index] == b'\'' => {
let backslash_escapes = index > 0
&& matches!(bytes[index - 1], b'e' | b'E')
&& (index < 2
|| !matches!(
bytes[index - 2],
b'a'..=b'z' | b'A'..=b'Z' | b'0'..=b'9' | b'_'
));
state = State::Single { backslash_escapes };
}
State::Normal if bytes[index] == b'"' => state = State::Double,
State::Normal if bytes[index] == b'$' => {
let start = index + 1;
let mut end = start;
while bytes.get(end).is_some_and(u8::is_ascii_digit) {
end += 1;
}
if end > start
&& let Ok(number) = sql[start..end].parse::<usize>()
{
count = count.max(number);
index = end - 1;
} else if let Some(delimiter_end) = dollar_quote_end(bytes, index) {
let delimiter = &bytes[index..delimiter_end];
let body_start = delimiter_end;
if let Some(offset) = find_bytes(&bytes[body_start..], delimiter) {
index = body_start + offset + delimiter.len() - 1;
} else {
index = bytes.len();
}
}
}
State::Single {
backslash_escapes: true,
} if bytes[index] == b'\\' => {
index += usize::from(index + 1 < bytes.len());
}
State::Single { .. } if bytes[index] == b'\'' => {
if bytes.get(index + 1) == Some(&b'\'') {
index += 1;
} else {
state = State::Normal;
}
}
State::Double if bytes[index] == b'"' => {
if bytes.get(index + 1) == Some(&b'"') {
index += 1;
} else {
state = State::Normal;
}
}
State::LineComment if matches!(bytes[index], b'\r' | b'\n') => {
state = State::Normal;
}
State::BlockComment(depth) if bytes[index..].starts_with(b"/*") => {
state = State::BlockComment(depth + 1);
index += 1;
}
State::BlockComment(depth) if bytes[index..].starts_with(b"*/") => {
state = if depth == 1 {
State::Normal
} else {
State::BlockComment(depth - 1)
};
index += 1;
}
State::Normal
| State::Single { .. }
| State::Double
| State::LineComment
| State::BlockComment(_) => {}
}
index += 1;
}
count
}
fn dollar_quote_end(bytes: &[u8], start: usize) -> Option<usize> {
let mut end = start + 1;
if bytes.get(end) == Some(&b'$') {
return Some(end + 1);
}
if !bytes
.get(end)
.is_some_and(|byte| byte.is_ascii_alphabetic() || *byte == b'_')
{
return None;
}
end += 1;
while bytes
.get(end)
.is_some_and(|byte| byte.is_ascii_alphanumeric() || *byte == b'_')
{
end += 1;
}
(bytes.get(end) == Some(&b'$')).then_some(end + 1)
}
fn find_bytes(haystack: &[u8], needle: &[u8]) -> Option<usize> {
haystack
.windows(needle.len())
.position(|window| window == needle)
}
fn classify(sql: &str) -> Statement {
let mut words = sql.split_whitespace();
let first = words.next().unwrap_or("").to_ascii_uppercase();
match first.as_str() {
"USE" => parse_use_target(sql).map_or(Statement::Query, Statement::Use),
"SHOW"
if words
.next()
.is_some_and(|word| word.eq_ignore_ascii_case("databases")) =>
{
Statement::ShowDatabases
}
"BEGIN" | "START" => Statement::Begin,
"COMMIT" | "END" => Statement::Commit,
"ROLLBACK" | "ABORT" => Statement::Rollback,
"SET" => Statement::Control("SET"),
"RESET" => Statement::Control("RESET"),
"DISCARD" => Statement::Control("DISCARD ALL"),
"INSERT" | "UPDATE" | "DELETE" => Statement::Mutation,
_ => Statement::Query,
}
}
fn parse_use_target(sql: &str) -> Option<String> {
let trimmed = sql.trim();
let rest = trimmed.get(3..)?.trim().trim_end_matches(';').trim();
if rest.is_empty() {
return None;
}
Some(unquote(rest))
}
fn unquote(value: &str) -> String {
if let Some(inner) = value
.strip_prefix('"')
.and_then(|rest| rest.strip_suffix('"'))
{
inner.replace("\"\"", "\"")
} else if let Some(inner) = value
.strip_prefix('\'')
.and_then(|rest| rest.strip_suffix('\''))
{
inner.replace("''", "'")
} else {
value.split_whitespace().next().unwrap_or("").to_owned()
}
}
fn first_keyword(sql: &str) -> &str {
sql.split_whitespace().next().unwrap_or("")
}
fn split_statements(input: &str) -> Vec<String> {
#[derive(Clone, Copy, Eq, PartialEq)]
enum State {
Normal,
SingleQuote,
DoubleQuote,
LineComment,
BlockComment,
}
let bytes = input.as_bytes();
let mut state = State::Normal;
let mut statements = Vec::new();
let mut start = 0;
let mut index = 0;
while index < bytes.len() {
let byte = bytes[index];
let next = bytes.get(index + 1).copied();
match state {
State::Normal => match (byte, next) {
(b'\'', _) => state = State::SingleQuote,
(b'"', _) => state = State::DoubleQuote,
(b'-', Some(b'-')) => {
state = State::LineComment;
index += 1;
}
(b'/', Some(b'*')) => {
state = State::BlockComment;
index += 1;
}
(b';', _) => {
let statement = input[start..index].trim();
if has_sql_content(statement) {
statements.push(statement.to_owned());
}
start = index + 1;
}
_ => {}
},
State::SingleQuote => {
if byte == b'\'' {
if next == Some(b'\'') {
index += 1;
} else {
state = State::Normal;
}
}
}
State::DoubleQuote => {
if byte == b'"' {
if next == Some(b'"') {
index += 1;
} else {
state = State::Normal;
}
}
}
State::LineComment if byte == b'\n' => state = State::Normal,
State::BlockComment if byte == b'*' && next == Some(b'/') => {
state = State::Normal;
index += 1;
}
State::LineComment | State::BlockComment => {}
}
index += 1;
}
let remainder = input[start..].trim();
if has_sql_content(remainder) {
statements.push(remainder.to_owned());
}
statements
}
fn has_sql_content(segment: &str) -> bool {
let bytes = segment.as_bytes();
let mut index = 0;
while index < bytes.len() {
let byte = bytes[index];
match (byte, bytes.get(index + 1).copied()) {
(b'-', Some(b'-')) => {
index += 2;
while index < bytes.len() && bytes[index] != b'\n' {
index += 1;
}
}
(b'/', Some(b'*')) => {
index += 2;
while index < bytes.len()
&& !(bytes[index] == b'*' && bytes.get(index + 1) == Some(&b'/'))
{
index += 1;
}
index += 2;
}
_ if byte.is_ascii_whitespace() => index += 1,
_ => return true,
}
}
false
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn statement_splitter_handles_quotes_and_trailing_statement() {
let statements = split_statements("SELECT ';'; SELECT \"a;b\"; SELECT 3");
assert_eq!(statements, vec!["SELECT ';'", "SELECT \"a;b\"", "SELECT 3"]);
}
#[test]
fn empty_query_splits_to_nothing() {
assert!(split_statements(" ; -- comment\n").is_empty());
}
#[test]
fn control_statements_are_recognized_case_insensitively() {
assert!(matches!(classify("begin"), Statement::Begin));
assert!(matches!(
classify(" SET client_encoding TO 'UTF8'"),
Statement::Control("SET")
));
assert!(matches!(classify("COMMIT"), Statement::Commit));
assert!(matches!(classify("SELECT 1"), Statement::Query));
}
#[test]
fn use_and_show_are_recognized() {
assert!(matches!(
classify("show databases"),
Statement::ShowDatabases
));
match classify("USE \"my-db\"") {
Statement::Use(name) => assert_eq!(name, "my-db"),
_ => panic!("expected USE"),
}
match classify("use people;") {
Statement::Use(name) => assert_eq!(name, "people"),
_ => panic!("expected USE"),
}
assert!(matches!(classify("USE"), Statement::Query));
}
#[test]
fn command_tag_counts_selects_and_names_explains() {
assert_eq!(command_tag("SELECT * FROM t", 7), "SELECT 7");
assert_eq!(command_tag("EXPLAIN SELECT 1", 3), "EXPLAIN");
}
#[test]
fn placeholder_count_skips_every_postgres_quoting_form() {
let sql = r#"
SELECT $2, '$99', E'escaped \' $98', "$97",
$$ body $96 $$, $tag$ body $95 $tag$
-- $94
/* $93 /* nested $92 */ still comment */
WHERE value = $7
"#;
assert_eq!(placeholder_count(sql), 7);
}
}