use crate::context::DjogiContext;
use crate::pg::connection::PgConnection;
use crate::pg::decode::{FromPgRow, try_get_scalar};
use crate::pg::pool::{ClientFuture, DjogiPool};
use crate::query::stream::{DEFAULT_FETCH_SIZE, RawCursorStream, build_raw_stream};
use crate::{DbError, DjogiError};
use postgres_types::{FromSql, ToSql};
use tokio_postgres::Row;
fn reject_transaction_backed_sql(ctx: &mut DjogiContext, sql: &str) -> Result<(), DjogiError> {
if let Some(err) = ctx.transaction_poison_error() {
return Err(err);
}
if ctx.conn().is_some()
&& let Some(refusal) = classify_transaction_backed_refusal(sql)
{
return Err(refusal.into_error());
}
Ok(())
}
fn reject_transaction_backed_sql_batch(
ctx: &mut DjogiContext,
sql: &str,
) -> Result<(), DjogiError> {
if let Some(err) = ctx.transaction_poison_error() {
return Err(err);
}
if ctx.conn().is_some()
&& let Some(refusal) = classify_raw_ddl_transaction_backed_refusal(sql)
{
return Err(refusal.into_error());
}
Ok(())
}
pub(crate) async fn guarded_batch_execute(
ctx: &mut DjogiContext,
sql: &str,
) -> Result<(), DjogiError> {
reject_transaction_backed_sql_batch(ctx, sql)?;
ctx.batch_execute(sql).await
}
fn classify_transaction_session_statement(sql: &str) -> Option<&'static str> {
let (keyword, next_idx) = parse_keyword(sql, 0)?;
if keyword.eq_ignore_ascii_case("SET") {
let second = parse_keyword(sql, next_idx).map(|(word, _)| word);
return match second {
Some(word)
if word.eq_ignore_ascii_case("LOCAL")
|| word.eq_ignore_ascii_case("CONSTRAINTS")
|| word.eq_ignore_ascii_case("TRANSACTION") =>
{
None
}
_ => Some("SET"),
};
}
[
"RESET",
"DISCARD",
"LISTEN",
"UNLISTEN",
"PREPARE",
"DEALLOCATE",
]
.into_iter()
.find(|statement| keyword.eq_ignore_ascii_case(statement))
}
#[allow(dead_code)] fn classify_raw_ddl_transaction_session_statement(sql: &str) -> Option<&'static str> {
let bytes = sql.as_bytes();
let mut statement_start = 0usize;
let mut idx = 0usize;
let mut block_comment_depth = 0usize;
let mut in_line_comment = false;
let mut in_single_quote = false;
let mut in_double_quote = false;
let mut dollar_quote: Option<String> = None;
while idx < bytes.len() {
if let Some(delimiter) = dollar_quote.as_deref() {
if bytes[idx..].starts_with(delimiter.as_bytes()) {
idx += delimiter.len();
dollar_quote = None;
} else {
idx += 1;
}
continue;
}
if in_line_comment {
if bytes[idx] == b'\n' {
in_line_comment = false;
}
idx += 1;
continue;
}
if block_comment_depth > 0 {
if bytes.get(idx) == Some(&b'/') && bytes.get(idx + 1) == Some(&b'*') {
block_comment_depth += 1;
idx += 2;
} else if bytes.get(idx) == Some(&b'*') && bytes.get(idx + 1) == Some(&b'/') {
block_comment_depth -= 1;
idx += 2;
} else {
idx += 1;
}
continue;
}
if in_single_quote {
if bytes[idx] == b'\'' {
if bytes.get(idx + 1) == Some(&b'\'') {
idx += 2;
} else {
in_single_quote = false;
idx += 1;
}
} else {
idx += 1;
}
continue;
}
if in_double_quote {
if bytes[idx] == b'"' {
if bytes.get(idx + 1) == Some(&b'"') {
idx += 2;
} else {
in_double_quote = false;
idx += 1;
}
} else {
idx += 1;
}
continue;
}
match bytes[idx] {
b';' => {
if let Some(statement) =
classify_transaction_session_statement(&sql[statement_start..idx])
{
return Some(statement);
}
statement_start = idx + 1;
idx += 1;
}
b'\'' => {
in_single_quote = true;
idx += 1;
}
b'"' => {
in_double_quote = true;
idx += 1;
}
b'-' if bytes.get(idx + 1) == Some(&b'-') => {
in_line_comment = true;
idx += 2;
}
b'/' if bytes.get(idx + 1) == Some(&b'*') => {
block_comment_depth = 1;
idx += 2;
}
b'$' => {
if let Some(end_idx) = parse_dollar_quote_delimiter_end(sql, idx) {
dollar_quote = Some(sql[idx..end_idx].to_owned());
idx = end_idx;
} else {
idx += 1;
}
}
_ => {
idx += 1;
}
}
}
classify_transaction_session_statement(&sql[statement_start..])
}
fn classify_transaction_control_statement(sql: &str) -> Option<&'static str> {
let (first, after_first) = parse_keyword(sql, 0)?;
if first.eq_ignore_ascii_case("START") {
return match parse_keyword(sql, after_first) {
Some((second, _)) if second.eq_ignore_ascii_case("TRANSACTION") => {
Some("START TRANSACTION")
}
_ => None,
};
}
if first.eq_ignore_ascii_case("ROLLBACK") {
let second = parse_keyword(sql, after_first);
return match second {
Some((w, _))
if w.eq_ignore_ascii_case("WORK") || w.eq_ignore_ascii_case("TRANSACTION") =>
{
Some("ROLLBACK")
}
Some((w, _)) if w.eq_ignore_ascii_case("TO") => Some("ROLLBACK"),
_ => Some("ROLLBACK"),
};
}
if first.eq_ignore_ascii_case("RELEASE") {
return Some("RELEASE");
}
["BEGIN", "COMMIT", "END", "ABORT", "SAVEPOINT"]
.into_iter()
.find(|s| first.eq_ignore_ascii_case(s))
}
#[derive(Debug, PartialEq)]
pub(crate) enum TransactionBackedRawSqlRefusal {
SessionStatement(&'static str),
TransactionControl(&'static str),
}
impl TransactionBackedRawSqlRefusal {
pub(crate) fn into_error(self) -> DjogiError {
match self {
Self::SessionStatement(s) => {
DjogiError::SessionStatementDisallowedInTransaction { statement: s }
}
Self::TransactionControl(s) => {
DjogiError::RawTransactionControlDisallowedInTransaction { statement: s }
}
}
}
}
fn classify_transaction_backed_refusal(sql: &str) -> Option<TransactionBackedRawSqlRefusal> {
if let Some(s) = classify_transaction_control_statement(sql) {
return Some(TransactionBackedRawSqlRefusal::TransactionControl(s));
}
if let Some(s) = classify_transaction_session_statement(sql) {
return Some(TransactionBackedRawSqlRefusal::SessionStatement(s));
}
None
}
fn classify_raw_ddl_transaction_backed_refusal(
sql: &str,
) -> Option<TransactionBackedRawSqlRefusal> {
let bytes = sql.as_bytes();
let mut statement_start = 0usize;
let mut idx = 0usize;
let mut block_comment_depth = 0usize;
let mut in_line_comment = false;
let mut in_single_quote = false;
let mut in_double_quote = false;
let mut dollar_quote: Option<String> = None;
while idx < bytes.len() {
if let Some(delimiter) = dollar_quote.as_deref() {
if bytes[idx..].starts_with(delimiter.as_bytes()) {
idx += delimiter.len();
dollar_quote = None;
} else {
idx += 1;
}
continue;
}
if in_line_comment {
if bytes[idx] == b'\n' {
in_line_comment = false;
}
idx += 1;
continue;
}
if block_comment_depth > 0 {
if bytes.get(idx) == Some(&b'/') && bytes.get(idx + 1) == Some(&b'*') {
block_comment_depth += 1;
idx += 2;
} else if bytes.get(idx) == Some(&b'*') && bytes.get(idx + 1) == Some(&b'/') {
block_comment_depth -= 1;
idx += 2;
} else {
idx += 1;
}
continue;
}
if in_single_quote {
if bytes[idx] == b'\'' {
if bytes.get(idx + 1) == Some(&b'\'') {
idx += 2;
} else {
in_single_quote = false;
idx += 1;
}
} else {
idx += 1;
}
continue;
}
if in_double_quote {
if bytes[idx] == b'"' {
if bytes.get(idx + 1) == Some(&b'"') {
idx += 2;
} else {
in_double_quote = false;
idx += 1;
}
} else {
idx += 1;
}
continue;
}
match bytes[idx] {
b';' => {
if let Some(refusal) =
classify_transaction_backed_refusal(&sql[statement_start..idx])
{
return Some(refusal);
}
statement_start = idx + 1;
idx += 1;
}
b'\'' => {
in_single_quote = true;
idx += 1;
}
b'"' => {
in_double_quote = true;
idx += 1;
}
b'-' if bytes.get(idx + 1) == Some(&b'-') => {
in_line_comment = true;
idx += 2;
}
b'/' if bytes.get(idx + 1) == Some(&b'*') => {
block_comment_depth = 1;
idx += 2;
}
b'$' => {
if let Some(end_idx) = parse_dollar_quote_delimiter_end(sql, idx) {
dollar_quote = Some(sql[idx..end_idx].to_owned());
idx = end_idx;
} else {
idx += 1;
}
}
_ => {
idx += 1;
}
}
}
classify_transaction_backed_refusal(&sql[statement_start..])
}
fn parse_keyword(sql: &str, start_idx: usize) -> Option<(&str, usize)> {
let bytes = sql.as_bytes();
let mut idx = skip_sql_trivia(sql, start_idx);
if idx >= bytes.len() || !bytes[idx].is_ascii_alphabetic() {
return None;
}
let start = idx;
idx += 1;
while idx < bytes.len() && (bytes[idx].is_ascii_alphanumeric() || bytes[idx] == b'_') {
idx += 1;
}
Some((&sql[start..idx], idx))
}
fn skip_sql_trivia(sql: &str, start_idx: usize) -> usize {
let bytes = sql.as_bytes();
let mut idx = start_idx;
loop {
while idx < bytes.len() && bytes[idx].is_ascii_whitespace() {
idx += 1;
}
if bytes.get(idx) == Some(&b'-') && bytes.get(idx + 1) == Some(&b'-') {
idx += 2;
while idx < bytes.len() && bytes[idx] != b'\n' {
idx += 1;
}
continue;
}
if bytes.get(idx) == Some(&b'/') && bytes.get(idx + 1) == Some(&b'*') {
idx += 2;
let mut depth = 1usize;
while idx < bytes.len() && depth > 0 {
if bytes.get(idx) == Some(&b'/') && bytes.get(idx + 1) == Some(&b'*') {
depth += 1;
idx += 2;
} else if bytes.get(idx) == Some(&b'*') && bytes.get(idx + 1) == Some(&b'/') {
depth -= 1;
idx += 2;
} else {
idx += 1;
}
}
continue;
}
return idx;
}
}
fn parse_dollar_quote_delimiter_end(sql: &str, start_idx: usize) -> Option<usize> {
let bytes = sql.as_bytes();
if bytes.get(start_idx) != Some(&b'$') {
return None;
}
let mut idx = start_idx + 1;
while idx < bytes.len() {
match bytes[idx] {
b'$' => return Some(idx + 1),
b'a'..=b'z' | b'A'..=b'Z' | b'0'..=b'9' | b'_' => idx += 1,
_ => return None,
}
}
None
}
mod sealed {
pub trait Sealed {}
impl Sealed for crate::context::DjogiContext {}
impl Sealed for crate::pg::pool::DjogiPool {}
}
#[doc(hidden)]
#[trait_variant::make(RawAccessExt: Send)]
pub trait RawAccessExtBase: sealed::Sealed {
async fn raw_query<T: FromPgRow>(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<Vec<T>, DjogiError>;
async fn raw_rows(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<Vec<Row>, DjogiError>;
async fn raw_fetch_one<T: FromPgRow>(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<T, DjogiError>;
async fn raw_scalar<T>(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<T, DjogiError>
where
T: for<'row> FromSql<'row> + Send + 'static;
async fn raw_execute(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<u64, DjogiError>;
async fn raw_ddl(&mut self, sql: &str) -> Result<(), DjogiError>;
async fn raw_stream<'ctx>(
&'ctx mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<RawCursorStream<'ctx>, DjogiError>;
async fn raw_stream_with_fetch_size<'ctx>(
&'ctx mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
fetch_size: u32,
) -> Result<RawCursorStream<'ctx>, DjogiError>;
}
impl RawAccessExt for DjogiContext {
async fn raw_query<T: FromPgRow>(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<Vec<T>, DjogiError> {
reject_transaction_backed_sql(self, sql)?;
self.query_all_with(sql, params, |rows| {
rows.iter().map(T::from_pg_row).collect()
})
.await
}
async fn raw_rows(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<Vec<Row>, DjogiError> {
reject_transaction_backed_sql(self, sql)?;
self.__query_all_for_macros(sql, params).await
}
async fn raw_fetch_one<T: FromPgRow>(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<T, DjogiError> {
reject_transaction_backed_sql(self, sql)?;
self.query_opt_with(sql, params, |row_opt| {
let row = row_opt.ok_or_else(|| DjogiError::not_found("<raw>"))?;
T::from_pg_row(&row)
})
.await
}
async fn raw_scalar<T>(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<T, DjogiError>
where
T: for<'row> FromSql<'row> + Send + 'static,
{
reject_transaction_backed_sql(self, sql)?;
self.query_opt_with(sql, params, |row_opt| {
let row = row_opt.ok_or_else(|| DjogiError::not_found("<raw>"))?;
try_get_scalar(&row, 0)
})
.await
}
async fn raw_execute(
&mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<u64, DjogiError> {
reject_transaction_backed_sql(self, sql)?;
self.__execute_for_macros(sql, params).await
}
async fn raw_ddl(&mut self, sql: &str) -> Result<(), DjogiError> {
reject_transaction_backed_sql_batch(self, sql)?;
self.batch_execute(sql).await
}
async fn raw_stream<'ctx>(
&'ctx mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
) -> Result<RawCursorStream<'ctx>, DjogiError> {
build_raw_stream(self, sql, params, DEFAULT_FETCH_SIZE).await
}
async fn raw_stream_with_fetch_size<'ctx>(
&'ctx mut self,
sql: &str,
params: &[&(dyn ToSql + Sync)],
fetch_size: u32,
) -> Result<RawCursorStream<'ctx>, DjogiError> {
if fetch_size == 0 {
return Err(DjogiError::Validation(
"raw_stream fetch_size must be at least 1".to_owned(),
));
}
build_raw_stream(self, sql, params, fetch_size).await
}
}
#[doc(hidden)]
#[trait_variant::make(RawPoolAccessExt: Send)]
pub trait RawPoolAccessExtBase: sealed::Sealed {
fn raw_pool(&self) -> Option<&DjogiPool>;
fn raw_conn(&mut self) -> Option<&mut PgConnection>;
async fn raw_with_client<F, R>(&self, f: F) -> Result<R, DjogiError>
where
F: for<'client> FnOnce(&'client mut tokio_postgres::Client) -> ClientFuture<'client, R>
+ Send,
R: Send + 'static;
}
impl RawPoolAccessExt for DjogiContext {
fn raw_pool(&self) -> Option<&DjogiPool> {
self.pool()
}
fn raw_conn(&mut self) -> Option<&mut PgConnection> {
if self.is_transaction_poisoned() {
None
} else {
self.conn()
}
}
fn raw_with_client<F, R>(
&self,
f: F,
) -> impl std::future::Future<Output = Result<R, DjogiError>> + Send
where
F: for<'client> FnOnce(&'client mut tokio_postgres::Client) -> ClientFuture<'client, R>
+ Send,
R: Send + 'static,
{
let pool = self.pool().cloned();
async move {
match pool {
Some(pool) => pool.with_client(f).await,
None => Err(DjogiError::Db(DbError::other(
"raw_with_client requires a pool-backed DjogiContext",
))),
}
}
}
}
impl RawPoolAccessExt for DjogiPool {
fn raw_pool(&self) -> Option<&DjogiPool> {
Some(self)
}
fn raw_conn(&mut self) -> Option<&mut PgConnection> {
None
}
fn raw_with_client<F, R>(
&self,
f: F,
) -> impl std::future::Future<Output = Result<R, DjogiError>> + Send
where
F: for<'client> FnOnce(&'client mut tokio_postgres::Client) -> ClientFuture<'client, R>
+ Send,
R: Send + 'static,
{
let pool = self.clone();
async move { pool.with_client(f).await }
}
}
#[cfg(test)]
#[allow(dead_code)]
async fn _raw_stream_trait_canary<'ctx>(
ctx: &'ctx mut DjogiContext,
) -> Result<RawCursorStream<'ctx>, DjogiError> {
let params: &[&(dyn ToSql + Sync)] = &[];
<DjogiContext as RawAccessExt>::raw_stream(ctx, "SELECT 1", params).await
}
#[cfg(test)]
#[allow(dead_code)]
async fn _raw_stream_with_fetch_size_trait_canary<'ctx>(
ctx: &'ctx mut DjogiContext,
) -> Result<RawCursorStream<'ctx>, DjogiError> {
let params: &[&(dyn ToSql + Sync)] = &[];
<DjogiContext as RawAccessExt>::raw_stream_with_fetch_size(ctx, "SELECT 1", params, 1).await
}
#[cfg(test)]
mod tests {
use super::{
TransactionBackedRawSqlRefusal, classify_raw_ddl_transaction_backed_refusal,
classify_raw_ddl_transaction_session_statement, classify_transaction_backed_refusal,
classify_transaction_control_statement, classify_transaction_session_statement,
};
use crate::DjogiError;
#[test]
fn classify_transaction_session_statement_rejects_plain_set_after_leading_comments() {
let sql = " /* prelude ; */ -- line comment\n sEt search_path = public";
assert_eq!(classify_transaction_session_statement(sql), Some("SET"));
}
#[test]
fn classify_transaction_session_statement_allows_transaction_local_set_forms() {
assert_eq!(
classify_transaction_session_statement("SET LOCAL statement_timeout = '5s'"),
None
);
assert_eq!(
classify_transaction_session_statement("SET CONSTRAINTS ALL IMMEDIATE"),
None
);
assert_eq!(
classify_transaction_session_statement(
"SET TRANSACTION ISOLATION LEVEL READ COMMITTED"
),
None
);
}
#[test]
fn classify_transaction_session_statement_rejects_other_session_statement_heads() {
for (sql, expected) in [
("RESET ALL", "RESET"),
("discard all", "DISCARD"),
("LISTEN djogi_updates", "LISTEN"),
("unlisten *", "UNLISTEN"),
("PREPARE x AS SELECT 1", "PREPARE"),
("deallocate all", "DEALLOCATE"),
] {
assert_eq!(
classify_transaction_session_statement(sql),
Some(expected),
"expected {expected} to be rejected for {sql:?}"
);
}
}
#[test]
fn classify_raw_ddl_transaction_session_statement_ignores_semicolons_inside_bodies() {
let sql = r#"
DO $body$
BEGIN
PERFORM '; still inside the body';
PERFORM $$nested ; dollar quote$$;
END
$body$;
/* scanner must only inspect the real next statement */
LISTEN djogi_updates;
"#;
assert_eq!(
classify_raw_ddl_transaction_session_statement(sql),
Some("LISTEN")
);
}
#[test]
fn classify_raw_ddl_transaction_session_statement_handles_utf8_inside_dollar_quote() {
let sql = r#"
DO $body$
BEGIN
-- Unicode comment inside the body: bootstrap — extensions
PERFORM 1;
END
$body$;
CREATE TEMP TABLE djogi_282_classifier_utf8_ok (value integer);
"#;
assert_eq!(classify_raw_ddl_transaction_session_statement(sql), None);
}
#[test]
fn classify_raw_ddl_transaction_session_statement_allows_trivia_only_and_safe_batches() {
assert_eq!(
classify_raw_ddl_transaction_session_statement(
" /* nothing here */ \n -- still nothing\n"
),
None
);
let sql = r#"
DO $body$
BEGIN
PERFORM '; safe body';
END
$body$;
CREATE TEMP TABLE djogi_282_classifier_ok (value integer);
"#;
assert_eq!(classify_raw_ddl_transaction_session_statement(sql), None);
}
#[test]
fn classify_raw_ddl_transaction_session_statement_rejects_session_set_in_batch() {
let sql = r#"
CREATE TEMP TABLE djogi_282_classifier_set_rejected (value integer);
SET statement_timeout = '1ms';
"#;
assert_eq!(
classify_raw_ddl_transaction_session_statement(sql),
Some("SET")
);
}
#[test]
fn classify_transaction_control_statement_detects_all_nine_forms() {
assert_eq!(
classify_transaction_control_statement("BEGIN"),
Some("BEGIN")
);
assert_eq!(
classify_transaction_control_statement("START TRANSACTION"),
Some("START TRANSACTION")
);
assert_eq!(
classify_transaction_control_statement("COMMIT"),
Some("COMMIT")
);
assert_eq!(
classify_transaction_control_statement("ROLLBACK"),
Some("ROLLBACK")
);
assert_eq!(classify_transaction_control_statement("END"), Some("END"));
assert_eq!(
classify_transaction_control_statement("ABORT"),
Some("ABORT")
);
assert_eq!(
classify_transaction_control_statement("SAVEPOINT my_sp"),
Some("SAVEPOINT")
);
assert_eq!(
classify_transaction_control_statement("RELEASE SAVEPOINT my_sp"),
Some("RELEASE")
);
assert_eq!(
classify_transaction_control_statement("RELEASE"),
Some("RELEASE")
);
assert_eq!(
classify_transaction_control_statement("ROLLBACK TO my_sp"),
Some("ROLLBACK")
);
assert_eq!(
classify_transaction_control_statement("ROLLBACK WORK TO my_sp"),
Some("ROLLBACK")
);
assert_eq!(
classify_transaction_control_statement("ROLLBACK TRANSACTION TO my_sp"),
Some("ROLLBACK")
);
assert_eq!(
classify_transaction_control_statement("ROLLBACK WORK"),
Some("ROLLBACK")
);
assert_eq!(
classify_transaction_control_statement("ROLLBACK TRANSACTION"),
Some("ROLLBACK")
);
assert_eq!(
classify_transaction_control_statement("COMMIT WORK"),
Some("COMMIT")
);
assert_eq!(
classify_transaction_control_statement("COMMIT TRANSACTION"),
Some("COMMIT")
);
assert_eq!(
classify_transaction_control_statement("END WORK"),
Some("END")
);
assert_eq!(
classify_transaction_control_statement("END TRANSACTION"),
Some("END")
);
}
#[test]
fn classify_transaction_control_statement_is_case_insensitive() {
for sql in ["commit", "CoMmIt", "COMMIT"] {
assert_eq!(
classify_transaction_control_statement(sql),
Some("COMMIT"),
"expected COMMIT for {sql:?}"
);
}
for sql in ["begin", "BeGiN", "BEGIN"] {
assert_eq!(
classify_transaction_control_statement(sql),
Some("BEGIN"),
"expected BEGIN for {sql:?}"
);
}
for sql in [
"start transaction",
"START TRANSACTION",
"Start Transaction",
] {
assert_eq!(
classify_transaction_control_statement(sql),
Some("START TRANSACTION"),
"expected START TRANSACTION for {sql:?}"
);
}
}
#[test]
fn classify_transaction_control_statement_handles_leading_trivia() {
assert_eq!(
classify_transaction_control_statement(" COMMIT"),
Some("COMMIT")
);
assert_eq!(
classify_transaction_control_statement(" \n \t rollback"),
Some("ROLLBACK")
);
assert_eq!(
classify_transaction_control_statement("-- line comment\nCOMMIT"),
Some("COMMIT")
);
assert_eq!(
classify_transaction_control_statement("/* block */ BEGIN"),
Some("BEGIN")
);
}
#[test]
fn classify_transaction_control_statement_returns_none_for_non_transaction_sql() {
for sql in [
"SELECT 1",
"INSERT INTO users (name) VALUES ('test')",
"UPDATE posts SET title = 'x'",
"DELETE FROM comments WHERE id = 1",
"CREATE TABLE foo (id bigint)",
"SET LOCAL statement_timeout = '5s'",
"SET CONSTRAINTS ALL IMMEDIATE",
"SET TRANSACTION ISOLATION LEVEL READ COMMITTED",
] {
assert_eq!(
classify_transaction_control_statement(sql),
None,
"expected None for non-transaction SQL: {sql:?}"
);
}
}
#[test]
fn classify_transaction_backed_refusal_prioritizes_transaction_control_over_session() {
let refusal = classify_transaction_backed_refusal("COMMIT").expect("expected refusal");
match refusal {
TransactionBackedRawSqlRefusal::TransactionControl(s) => {
assert_eq!(s, "COMMIT");
}
_ => panic!("expected TransactionControl(COMMIT), got {:?}", refusal),
}
}
#[test]
fn classify_transaction_backed_refusal_wraps_session_statements() {
let refusal = classify_transaction_backed_refusal("RESET ALL").expect("expected refusal");
match refusal {
TransactionBackedRawSqlRefusal::SessionStatement(s) => {
assert_eq!(s, "RESET");
}
_ => panic!("expected SessionStatement(RESET), got {:?}", refusal),
}
}
#[test]
fn classify_transaction_backed_refusal_into_error_produces_correct_variant() {
let refusal = classify_transaction_backed_refusal("COMMIT").expect("expected refusal");
let err = refusal.into_error();
match err {
DjogiError::RawTransactionControlDisallowedInTransaction { statement } => {
assert_eq!(statement, "COMMIT");
}
_ => panic!(
"expected RawTransactionControlDisallowedInTransaction, got {:?}",
err
),
}
let refusal = classify_transaction_backed_refusal("LISTEN foo").expect("expected refusal");
let err = refusal.into_error();
match err {
DjogiError::SessionStatementDisallowedInTransaction { statement } => {
assert_eq!(statement, "LISTEN");
}
_ => panic!(
"expected SessionStatementDisallowedInTransaction, got {:?}",
err
),
}
}
#[test]
fn classify_raw_ddl_batch_ignores_transaction_keywords_in_dollar_quoted_body() {
let sql = r#"
DO $body$
BEGIN
PERFORM 'COMMIT should be ignored here';
PERFORM $$nested ROLLBACK$$;
END
$body$;
SELECT 1;
"#;
assert_eq!(classify_raw_ddl_transaction_backed_refusal(sql), None);
}
#[test]
fn classify_raw_ddl_batch_detects_transaction_control_after_safe_ddl() {
let sql = r#"
CREATE TEMP TABLE foo (value integer);
COMMIT;
"#;
match classify_raw_ddl_transaction_backed_refusal(sql) {
Some(TransactionBackedRawSqlRefusal::TransactionControl(s)) => {
assert_eq!(s, "COMMIT");
}
_ => panic!("expected TransactionControl(COMMIT)"),
}
}
#[test]
fn classify_raw_ddl_batch_detects_session_statement_after_safe_ddl() {
let sql = r#"
CREATE TEMP TABLE foo (value integer);
RESET ALL;
"#;
match classify_raw_ddl_transaction_backed_refusal(sql) {
Some(TransactionBackedRawSqlRefusal::SessionStatement(s)) => {
assert_eq!(s, "RESET");
}
_ => panic!("expected SessionStatement(RESET)"),
}
}
#[test]
fn classify_raw_ddl_batch_allows_trivia_only_and_safe_batches() {
assert_eq!(
classify_raw_ddl_transaction_backed_refusal(
" /* nothing here */ \n -- still nothing\n"
),
None
);
let sql = r#"
DO $body$
BEGIN
PERFORM '; safe body';
END
$body$;
CREATE TEMP TABLE djogi_306_classifier_ok (value integer);
"#;
assert_eq!(classify_raw_ddl_transaction_backed_refusal(sql), None);
}
#[test]
fn classify_transaction_control_statement_start_without_transaction_is_none() {
assert_eq!(classify_transaction_control_statement("START"), None);
assert_eq!(classify_transaction_control_statement("START ALL"), None);
}
#[test]
fn classify_transaction_backed_refusal_returns_none_for_safe_sql() {
assert_eq!(classify_transaction_backed_refusal("SELECT 1"), None);
assert_eq!(
classify_transaction_backed_refusal("INSERT INTO t VALUES (1)"),
None
);
}
}