diff --git a/postgres/src/client.rs b/postgres/src/client.rs index 2dce6d94..ffa3169e 100644 --- a/postgres/src/client.rs +++ b/postgres/src/client.rs @@ -6,7 +6,7 @@ use tokio_postgres::{MakeTlsMode, Socket, TlsMode}; #[cfg(feature = "runtime")] use crate::Builder; -use crate::Statement; +use crate::{Statement, Transaction}; pub struct Client(tokio_postgres::Client); @@ -51,6 +51,11 @@ impl Client { pub fn batch_execute(&mut self, query: &str) -> Result<(), Error> { self.0.batch_execute(query).wait() } + + pub fn transaction(&mut self) -> Result, Error> { + self.batch_execute("BEGIN")?; + Ok(Transaction::new(self)) + } } impl From for Client { diff --git a/postgres/src/lib.rs b/postgres/src/lib.rs index 1072ea7c..acd88f4d 100644 --- a/postgres/src/lib.rs +++ b/postgres/src/lib.rs @@ -7,11 +7,13 @@ use tokio::runtime::{self, Runtime}; mod builder; mod client; mod statement; +mod transaction; #[cfg(feature = "runtime")] pub use crate::builder::*; pub use crate::client::*; pub use crate::statement::*; +pub use crate::transaction::*; #[cfg(feature = "runtime")] lazy_static! { diff --git a/postgres/src/transaction.rs b/postgres/src/transaction.rs new file mode 100644 index 00000000..56d2dcd3 --- /dev/null +++ b/postgres/src/transaction.rs @@ -0,0 +1,64 @@ +use tokio_postgres::types::{ToSql, Type}; +use tokio_postgres::{Error, Row}; + +use crate::{Client, Statement}; + +pub struct Transaction<'a> { + client: &'a mut Client, + done: bool, +} + +impl<'a> Drop for Transaction<'a> { + fn drop(&mut self) { + if !self.done { + let _ = self.rollback_inner(); + } + } +} + +impl<'a> Transaction<'a> { + pub(crate) fn new(client: &'a mut Client) -> Transaction<'a> { + Transaction { + client, + done: false, + } + } + + pub fn commit(mut self) -> Result<(), Error> { + self.done = true; + self.client.batch_execute("COMMIT") + } + + pub fn rollback(mut self) -> Result<(), Error> { + self.done = true; + self.rollback_inner() + } + + fn rollback_inner(&mut self) -> Result<(), Error> { + self.client.batch_execute("ROLLBACK") + } + + pub fn prepare(&mut self, query: &str) -> Result { + self.client.prepare(query) + } + + pub fn prepare_typed(&mut self, query: &str, types: &[Type]) -> Result { + self.client.prepare_typed(query, types) + } + + pub fn execute(&mut self, statement: &Statement, params: &[&dyn ToSql]) -> Result { + self.client.execute(statement, params) + } + + pub fn query( + &mut self, + statement: &Statement, + params: &[&dyn ToSql], + ) -> Result, Error> { + self.client.query(statement, params) + } + + pub fn batch_execute(&mut self, query: &str) -> Result<(), Error> { + self.client.batch_execute(query) + } +}