use super::{PgConnection, PgError, PgResult};
fn quote_savepoint_name(name: &str) -> PgResult<String> {
if name.is_empty() {
return Err(PgError::Query("savepoint name is empty".to_string()));
}
if name.contains('\0') {
return Err(PgError::Query(
"savepoint name contains NUL byte".to_string(),
));
}
Ok(format!("\"{}\"", name.replace('"', "\"\"")))
}
pub(crate) fn commit_outcome(tags: &[String], operation: &str) -> PgResult<()> {
match tags.first().map(String::as_str) {
Some("COMMIT") => Ok(()),
Some(tag) => Err(PgError::Query(format!(
"{operation}: the server answered COMMIT with {tag}; \
the transaction had failed and none of its writes were kept"
))),
None => Err(PgError::Query(format!(
"{operation}: the server answered COMMIT without a command tag; \
whether the transaction's writes were kept is unknown"
))),
}
}
impl PgConnection {
pub async fn begin_transaction(&mut self) -> PgResult<()> {
self.execute_simple("BEGIN").await
}
pub async fn commit(&mut self) -> PgResult<()> {
let tags = self.execute_simple_tags("COMMIT").await?;
commit_outcome(&tags, "COMMIT")
}
pub async fn rollback(&mut self) -> PgResult<()> {
self.execute_simple("ROLLBACK").await
}
pub async fn savepoint(&mut self, name: &str) -> PgResult<()> {
self.execute_simple(&format!("SAVEPOINT {}", quote_savepoint_name(name)?))
.await
}
pub async fn rollback_to(&mut self, name: &str) -> PgResult<()> {
self.execute_simple(&format!(
"ROLLBACK TO SAVEPOINT {}",
quote_savepoint_name(name)?
))
.await
}
pub async fn release_savepoint(&mut self, name: &str) -> PgResult<()> {
self.execute_simple(&format!(
"RELEASE SAVEPOINT {}",
quote_savepoint_name(name)?
))
.await
}
}
#[cfg(test)]
mod tests {
use super::{commit_outcome, quote_savepoint_name};
fn tags(tags: &[&str]) -> Vec<String> {
tags.iter().map(|tag| tag.to_string()).collect()
}
#[test]
fn commit_outcome_keeps_a_commit_tag() {
assert!(commit_outcome(&tags(&["COMMIT"]), "COMMIT").is_ok());
assert!(commit_outcome(&tags(&["COMMIT", "CLOSE CURSOR ALL", "SET"]), "COMMIT").is_ok());
}
#[test]
fn commit_outcome_reports_a_commit_answered_rollback() {
let err = commit_outcome(
&tags(&["ROLLBACK", "CLOSE CURSOR ALL"]),
"pool release COMMIT",
)
.expect_err("a rolled-back COMMIT is an error");
let text = err.to_string();
assert!(text.contains("pool release COMMIT"), "{text}");
assert!(text.contains("ROLLBACK"), "{text}");
}
#[test]
fn commit_outcome_reports_a_commit_without_a_tag() {
assert!(commit_outcome(&[], "COMMIT").is_err());
}
#[test]
fn quote_savepoint_name_escapes_quotes() {
assert_eq!(quote_savepoint_name("sp\"1").unwrap(), "\"sp\"\"1\"");
}
#[test]
fn quote_savepoint_name_rejects_empty_or_nul() {
assert!(quote_savepoint_name("").is_err());
assert!(quote_savepoint_name("sp\0shadow").is_err());
}
}