use crate::connection::Connection;
use crate::error::Result;
use crate::query::result::{ExecuteResult, QueryResult};
#[cfg(feature = "tracing")]
use crate::tracing_ext::TARGET_TRANSACTION;
pub mod options;
pub mod savepoint;
pub use options::{IsolationLevel, TransactionOptions};
pub use savepoint::Savepoint;
#[non_exhaustive]
pub struct Transaction<'a> {
pub(crate) conn: &'a mut Connection,
pub(crate) committed: bool,
pub(crate) savepoint_depth: u32,
}
impl<'a> Transaction<'a> {
pub(crate) fn new(conn: &'a mut Connection) -> Self {
Self {
conn,
committed: false,
savepoint_depth: 0,
}
}
#[must_use = "commit errors should be checked"]
pub async fn commit(mut self) -> Result<()> {
#[cfg(feature = "tracing")]
tracing::info!(target: TARGET_TRANSACTION, "COMMIT transaction");
self.conn.execute("COMMIT").await?;
self.committed = true;
Ok(())
}
#[must_use = "rollback errors should be checked"]
pub async fn rollback(mut self) -> Result<()> {
#[cfg(feature = "tracing")]
tracing::warn!(target: TARGET_TRANSACTION, "ROLLBACK transaction");
self.conn.execute("ROLLBACK").await?;
self.committed = true;
Ok(())
}
pub fn is_failed(&self) -> bool {
self.conn.transaction_status() == crate::protocol::TransactionStatus::Failed
}
#[must_use = "query errors should be checked"]
pub async fn query(&mut self, sql: &str) -> Result<QueryResult> {
self.conn.query(sql).await
}
#[must_use = "execute errors should be checked"]
pub async fn execute(&mut self, sql: &str) -> Result<ExecuteResult> {
self.conn.execute(sql).await
}
#[must_use = "query errors should be checked"]
pub async fn query_one(&mut self, sql: &str) -> Result<Option<crate::Row>> {
self.conn.query_one(sql).await
}
#[must_use = "query errors should be checked"]
pub async fn query_params(
&mut self,
sql: &str,
params: &[&dyn crate::types::ToSql],
) -> Result<QueryResult> {
self.conn.query_params(sql, params).await
}
#[must_use = "execute errors should be checked"]
pub async fn execute_params(
&mut self,
sql: &str,
params: &[&dyn crate::types::ToSql],
) -> Result<ExecuteResult> {
self.conn.execute_params(sql, params).await
}
#[must_use = "prepare errors should be checked"]
pub async fn prepare(&mut self, sql: &str) -> Result<crate::query::PreparedStatement> {
self.conn.prepare(sql).await
}
#[must_use = "savepoint errors should be checked"]
pub async fn savepoint(&mut self, name: &str) -> Result<Savepoint<'_, 'a>> {
let sql = format!("SAVEPOINT {}", quote_identifier(name));
self.conn.execute(&sql).await?;
self.savepoint_depth += 1;
Ok(Savepoint {
transaction: self,
name: name.to_string(),
released: false,
})
}
#[must_use = "copy errors should be checked"]
pub async fn copy_in(&mut self, sql: &str) -> Result<crate::CopyIn<'_>> {
self.conn.copy_in(sql).await
}
#[must_use = "copy errors should be checked"]
pub async fn copy_out(&mut self, sql: &str) -> Result<crate::CopyOut<'_>> {
self.conn.copy_out(sql).await
}
}
impl<'a> Drop for Transaction<'a> {
#[allow(clippy::needless_return)]
fn drop(&mut self) {
if self.committed || std::thread::panicking() {
return;
}
#[cfg(feature = "tracing")]
tracing::warn!(target: TARGET_TRANSACTION, "Transaction dropped without explicit commit/rollback");
}
}
pub(crate) fn quote_identifier(name: &str) -> String {
format!("\"{}\"", name.replace('"', "\"\""))
}
impl Connection {
#[must_use = "transaction errors should be checked"]
pub async fn transaction(&mut self) -> Result<Transaction<'_>> {
self.execute("BEGIN").await?;
#[cfg(feature = "tracing")]
tracing::info!(target: TARGET_TRANSACTION, "BEGIN transaction");
Ok(Transaction::new(self))
}
#[must_use = "transaction errors should be checked"]
pub async fn transaction_with(
&mut self,
options: &TransactionOptions,
) -> Result<Transaction<'_>> {
let sql = options.to_begin_sql();
self.execute(&sql).await?;
Ok(Transaction::new(self))
}
#[must_use = "transaction errors should be checked"]
pub async fn with_transaction<T, F>(&mut self, f: F) -> Result<T>
where
F: AsyncFnOnce(&mut Transaction<'_>) -> Result<T>,
{
let mut txn = self.transaction().await?;
match f(&mut txn).await {
Ok(val) => {
txn.commit().await?;
Ok(val)
}
Err(e) => {
let _ = txn.rollback().await;
Err(e)
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::auth::{Codec, ServerParams};
use crate::config::Config;
use crate::connection::ConnectionState;
use crate::error::Error;
use crate::protocol::TransactionStatus;
use crate::transport::{BufferedTransport, ClientTransport, MockTransport, PgTransport};
use std::collections::VecDeque;
fn make_connection(read_data: Vec<u8>) -> Connection {
let transport = PgTransport::Plain(BufferedTransport::new(ClientTransport::Mock(
MockTransport::new(read_data),
)));
Connection {
transport,
codec: Codec::new(),
server_params: ServerParams::default(),
state: ConnectionState::Idle,
config: Config::new(),
transaction_status: TransactionStatus::Idle,
notification_queue: VecDeque::new(),
notice_handler: None,
statement_counter: 0,
needs_recovery: false,
health: crate::reconnect::session::ConnectionHealth::new(),
session_state: crate::reconnect::session::SessionState::new(),
}
}
fn build_command_complete_msg(tag: &str) -> Vec<u8> {
let mut buf = vec![b'C'];
let mut body = Vec::new();
body.extend_from_slice(tag.as_bytes());
body.push(0);
let len = (body.len() + 4) as i32;
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(&body);
buf
}
fn build_ready_for_query(status: u8) -> Vec<u8> {
vec![b'Z', 0, 0, 0, 5, status]
}
fn build_error_response(msg: &str) -> Vec<u8> {
let mut buf = vec![b'E'];
let mut body = Vec::new();
body.push(b'S');
body.extend_from_slice(b"ERROR\0");
body.push(b'M');
body.extend_from_slice(msg.as_bytes());
body.push(0);
body.push(0);
let len = (body.len() + 4) as i32;
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(&body);
buf
}
fn build_row_description_msg(fields: &[(&str, u32)]) -> Vec<u8> {
let mut buf = vec![b'T'];
let mut body = Vec::new();
body.extend_from_slice(&(fields.len() as i16).to_be_bytes());
for (name, type_oid) in fields {
body.extend_from_slice(name.as_bytes());
body.push(0);
body.extend_from_slice(&0u32.to_be_bytes()); body.extend_from_slice(&0i16.to_be_bytes()); body.extend_from_slice(&type_oid.to_be_bytes()); body.extend_from_slice(&(-1i16).to_be_bytes()); body.extend_from_slice(&(-1i32).to_be_bytes()); body.extend_from_slice(&0i16.to_be_bytes()); }
let len = (body.len() + 4) as i32;
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(&body);
buf
}
fn build_data_row_msg(values: &[Option<&str>]) -> Vec<u8> {
let mut buf = vec![b'D'];
let mut body = Vec::new();
body.extend_from_slice(&(values.len() as i16).to_be_bytes());
for val in values {
match val {
Some(v) => {
let bytes = v.as_bytes();
body.extend_from_slice(&(bytes.len() as i32).to_be_bytes());
body.extend_from_slice(bytes);
}
None => {
body.extend_from_slice(&(-1i32).to_be_bytes());
}
}
}
let len = (body.len() + 4) as i32;
buf.extend_from_slice(&len.to_be_bytes());
buf.extend_from_slice(&body);
buf
}
#[test]
fn test_quote_identifier_basic() {
assert_eq!(quote_identifier("foo"), "\"foo\"");
}
#[test]
fn test_quote_identifier_with_quotes() {
assert_eq!(quote_identifier("foo\"bar"), "\"foo\"\"bar\"");
}
#[test]
fn test_quote_identifier_empty() {
assert_eq!(quote_identifier(""), "\"\"");
}
#[test]
fn test_transaction_options_default() {
let opts = TransactionOptions::new();
assert_eq!(opts.to_begin_sql(), "BEGIN");
}
#[test]
fn test_transaction_options_isolation() {
let opts = TransactionOptions::new().isolation_level(IsolationLevel::Serializable);
assert_eq!(opts.to_begin_sql(), "BEGIN ISOLATION LEVEL SERIALIZABLE");
}
#[test]
fn test_transaction_options_all() {
let opts = TransactionOptions::new()
.isolation_level(IsolationLevel::RepeatableRead)
.read_only(true)
.deferrable(true);
assert_eq!(
opts.to_begin_sql(),
"BEGIN ISOLATION LEVEL REPEATABLE READ READ ONLY DEFERRABLE"
);
}
#[test]
fn test_transaction_options_read_write() {
let opts = TransactionOptions::new().read_only(false);
assert_eq!(opts.to_begin_sql(), "BEGIN READ WRITE");
}
#[tokio::test]
async fn test_transaction_commit_mock() {
let mut data = Vec::new();
data.extend_from_slice(&build_command_complete_msg("BEGIN"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_command_complete_msg("COMMIT"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let txn = conn.transaction().await.unwrap();
assert!(!txn.committed);
txn.commit().await.unwrap();
assert_eq!(conn.transaction_status(), TransactionStatus::Idle);
}
#[tokio::test]
async fn test_transaction_rollback_mock() {
let mut data = Vec::new();
data.extend_from_slice(&build_command_complete_msg("BEGIN"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_command_complete_msg("ROLLBACK"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let txn = conn.transaction().await.unwrap();
assert!(!txn.committed);
txn.rollback().await.unwrap();
assert_eq!(conn.transaction_status(), TransactionStatus::Idle);
}
#[tokio::test]
async fn test_transaction_is_failed_mock() {
let mut data = Vec::new();
data.extend_from_slice(&build_command_complete_msg("BEGIN"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_error_response("syntax error"));
data.extend_from_slice(&build_ready_for_query(b'E'));
data.extend_from_slice(&build_command_complete_msg("ROLLBACK"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let mut txn = conn.transaction().await.unwrap();
assert!(!txn.is_failed());
let err = txn.execute("BAD SQL").await;
assert!(err.is_err());
assert!(txn.is_failed());
txn.rollback().await.unwrap();
assert_eq!(conn.transaction_status(), TransactionStatus::Idle);
}
#[tokio::test]
async fn test_transaction_query_delegation_mock() {
let mut data = Vec::new();
data.extend_from_slice(&build_command_complete_msg("BEGIN"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_row_description_msg(&[(
"val",
crate::types::INT4_OID,
)]));
data.extend_from_slice(&build_data_row_msg(&[Some("42")]));
data.extend_from_slice(&build_command_complete_msg("SELECT 1"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_command_complete_msg("COMMIT"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let mut txn = conn.transaction().await.unwrap();
let result = txn.query("SELECT 42").await.unwrap();
assert_eq!(result.len(), 1);
let v: i32 = result.rows()[0].get(0).unwrap();
assert_eq!(v, 42);
txn.commit().await.unwrap();
}
#[tokio::test]
async fn test_with_transaction_success_mock() {
let mut data = Vec::new();
data.extend_from_slice(&build_command_complete_msg("BEGIN"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_row_description_msg(&[(
"val",
crate::types::INT4_OID,
)]));
data.extend_from_slice(&build_data_row_msg(&[Some("42")]));
data.extend_from_slice(&build_command_complete_msg("SELECT 1"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_command_complete_msg("COMMIT"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let result = conn
.with_transaction(async |txn| {
let qr = txn.query("SELECT 42").await?;
let v: i32 = qr.rows()[0].get(0)?;
Ok(v)
})
.await
.unwrap();
assert_eq!(result, 42);
assert_eq!(conn.transaction_status(), TransactionStatus::Idle);
}
#[tokio::test]
async fn test_with_transaction_error_rolls_back_mock() {
let mut data = Vec::new();
data.extend_from_slice(&build_command_complete_msg("BEGIN"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_row_description_msg(&[(
"val",
crate::types::INT4_OID,
)]));
data.extend_from_slice(&build_data_row_msg(&[Some("42")]));
data.extend_from_slice(&build_command_complete_msg("SELECT 1"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_command_complete_msg("ROLLBACK"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let result = conn
.with_transaction(async |txn| {
let qr = txn.query("SELECT 42").await?;
let _v: i32 = qr.rows()[0].get(0)?;
Err::<i32, Error>(Error::Config("intentional failure".into()))
})
.await;
assert!(result.is_err());
assert_eq!(conn.transaction_status(), TransactionStatus::Idle);
}
#[tokio::test]
async fn test_transaction_savepoint_mock() {
let mut data = Vec::new();
data.extend_from_slice(&build_command_complete_msg("BEGIN"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_command_complete_msg("SAVEPOINT sp1"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_command_complete_msg("RELEASE SAVEPOINT sp1"));
data.extend_from_slice(&build_ready_for_query(b'T'));
data.extend_from_slice(&build_command_complete_msg("COMMIT"));
data.extend_from_slice(&build_ready_for_query(b'I'));
let mut conn = make_connection(data);
let mut txn = conn.transaction().await.unwrap();
assert_eq!(txn.savepoint_depth, 0);
let sp = txn.savepoint("sp1").await.unwrap();
sp.release().await.unwrap();
assert_eq!(txn.savepoint_depth, 0);
txn.commit().await.unwrap();
}
#[test]
fn test_transaction_drop_without_commit_does_not_panic() {
let mut conn = make_connection(Vec::new());
let txn = Transaction::new(&mut conn);
assert!(!txn.committed);
drop(txn); }
#[test]
fn test_transaction_drop_after_commit_is_noop() {
let mut conn = make_connection(Vec::new());
let txn = Transaction::new(&mut conn);
let mut txn = txn;
txn.committed = true;
drop(txn); }
}