prax_postgres/
connection.rs1use std::sync::Arc;
4
5use deadpool_postgres::Object;
6use tokio_postgres::Row;
7use tracing::{debug, trace};
8
9use prax_query::sql::is_valid_sql_identifier;
10
11use crate::error::{PgError, PgResult};
12use crate::statement::PreparedStatementCache;
13
14pub struct PgConnection {
16 client: Object,
17 statement_cache: Arc<PreparedStatementCache>,
18}
19
20impl PgConnection {
21 pub(crate) fn new(client: Object, statement_cache: Arc<PreparedStatementCache>) -> Self {
23 Self {
24 client,
25 statement_cache,
26 }
27 }
28
29 pub async fn query(
31 &self,
32 sql: &str,
33 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
34 ) -> PgResult<Vec<Row>> {
35 trace!(sql = %sql, "Executing query");
36
37 let stmt = self
39 .statement_cache
40 .get_or_prepare(&self.client, sql)
41 .await?;
42
43 let rows = self.client.query(&stmt, params).await?;
44 Ok(rows)
45 }
46
47 pub async fn query_one(
49 &self,
50 sql: &str,
51 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
52 ) -> PgResult<Row> {
53 trace!(sql = %sql, "Executing query_one");
54
55 let stmt = self
56 .statement_cache
57 .get_or_prepare(&self.client, sql)
58 .await?;
59
60 let row = self.client.query_one(&stmt, params).await?;
61 Ok(row)
62 }
63
64 pub async fn query_opt(
66 &self,
67 sql: &str,
68 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
69 ) -> PgResult<Option<Row>> {
70 trace!(sql = %sql, "Executing query_opt");
71
72 let stmt = self
73 .statement_cache
74 .get_or_prepare(&self.client, sql)
75 .await?;
76
77 let row = self.client.query_opt(&stmt, params).await?;
78 Ok(row)
79 }
80
81 pub async fn execute(
83 &self,
84 sql: &str,
85 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
86 ) -> PgResult<u64> {
87 trace!(sql = %sql, "Executing statement");
88
89 let stmt = self
90 .statement_cache
91 .get_or_prepare(&self.client, sql)
92 .await?;
93
94 let count = self.client.execute(&stmt, params).await?;
95 Ok(count)
96 }
97
98 pub async fn batch_execute(&self, sql: &str) -> PgResult<()> {
100 trace!(sql = %sql, "Executing batch");
101 self.client.batch_execute(sql).await?;
102 Ok(())
103 }
104
105 pub async fn transaction(&mut self) -> PgResult<PgTransaction<'_>> {
107 debug!("Beginning transaction");
108 let txn = self.client.transaction().await?;
109 Ok(PgTransaction {
110 txn,
111 statement_cache: self.statement_cache.clone(),
112 })
113 }
114
115 pub fn inner(&self) -> &Object {
119 &self.client
120 }
121
122 #[inline]
128 pub async fn query_cached(
129 &self,
130 sql: &str,
131 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
132 ) -> PgResult<Vec<Row>> {
133 self.query(sql, params).await
134 }
135
136 pub async fn query_raw(
141 &self,
142 sql: &str,
143 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
144 ) -> PgResult<Vec<Row>> {
145 trace!(sql = %sql, "Executing raw query (no statement cache)");
146 let rows = self.client.query(sql, params).await?;
147 Ok(rows)
148 }
149
150 pub async fn query_opt_raw(
152 &self,
153 sql: &str,
154 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
155 ) -> PgResult<Option<Row>> {
156 trace!(sql = %sql, "Executing raw query_opt (no statement cache)");
157 let row = self.client.query_opt(sql, params).await?;
158 Ok(row)
159 }
160}
161
162const MAX_SAVEPOINT_NAME_LEN: usize = 63;
164
165fn validate_savepoint_name(name: &str) -> PgResult<()> {
170 let valid = name.len() <= MAX_SAVEPOINT_NAME_LEN && is_valid_sql_identifier(name);
171 if !valid {
172 return Err(PgError::query(format!("invalid savepoint name: {name:?}")));
173 }
174 Ok(())
175}
176
177pub struct PgTransaction<'a> {
179 txn: deadpool_postgres::Transaction<'a>,
180 statement_cache: Arc<PreparedStatementCache>,
181}
182
183impl<'a> PgTransaction<'a> {
184 pub async fn query(
186 &self,
187 sql: &str,
188 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
189 ) -> PgResult<Vec<Row>> {
190 trace!(sql = %sql, "Executing query in transaction");
191
192 let stmt = self
193 .statement_cache
194 .get_or_prepare_in_txn(&self.txn, sql)
195 .await?;
196
197 let rows = self.txn.query(&stmt, params).await?;
198 Ok(rows)
199 }
200
201 pub async fn query_one(
203 &self,
204 sql: &str,
205 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
206 ) -> PgResult<Row> {
207 let stmt = self
208 .statement_cache
209 .get_or_prepare_in_txn(&self.txn, sql)
210 .await?;
211
212 let row = self.txn.query_one(&stmt, params).await?;
213 Ok(row)
214 }
215
216 pub async fn query_opt(
218 &self,
219 sql: &str,
220 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
221 ) -> PgResult<Option<Row>> {
222 let stmt = self
223 .statement_cache
224 .get_or_prepare_in_txn(&self.txn, sql)
225 .await?;
226
227 let row = self.txn.query_opt(&stmt, params).await?;
228 Ok(row)
229 }
230
231 pub async fn execute(
233 &self,
234 sql: &str,
235 params: &[&(dyn tokio_postgres::types::ToSql + Sync)],
236 ) -> PgResult<u64> {
237 let stmt = self
238 .statement_cache
239 .get_or_prepare_in_txn(&self.txn, sql)
240 .await?;
241
242 let count = self.txn.execute(&stmt, params).await?;
243 Ok(count)
244 }
245
246 pub async fn savepoint(&mut self, name: &str) -> PgResult<()> {
248 validate_savepoint_name(name)?;
249 debug!(name = %name, "Creating savepoint");
250 self.txn
251 .batch_execute(&format!("SAVEPOINT {}", name))
252 .await?;
253 Ok(())
254 }
255
256 pub async fn rollback_to(&mut self, name: &str) -> PgResult<()> {
258 validate_savepoint_name(name)?;
259 debug!(name = %name, "Rolling back to savepoint");
260 self.txn
261 .batch_execute(&format!("ROLLBACK TO SAVEPOINT {}", name))
262 .await?;
263 Ok(())
264 }
265
266 pub async fn release_savepoint(&mut self, name: &str) -> PgResult<()> {
268 validate_savepoint_name(name)?;
269 debug!(name = %name, "Releasing savepoint");
270 self.txn
271 .batch_execute(&format!("RELEASE SAVEPOINT {}", name))
272 .await?;
273 Ok(())
274 }
275
276 pub async fn commit(self) -> PgResult<()> {
278 debug!("Committing transaction");
279 self.txn.commit().await?;
280 Ok(())
281 }
282
283 pub async fn rollback(self) -> PgResult<()> {
285 debug!("Rolling back transaction");
286 self.txn.rollback().await?;
287 Ok(())
288 }
289}
290
291#[cfg(test)]
292mod tests {
293 use super::*;
294
295 #[test]
299 fn test_validate_savepoint_name_accepts_valid_names() {
300 assert!(validate_savepoint_name("sp1").is_ok());
301 assert!(validate_savepoint_name("my_savepoint").is_ok());
302 assert!(validate_savepoint_name("_private").is_ok());
303 assert!(validate_savepoint_name("SP_2").is_ok());
304 assert!(validate_savepoint_name("a").is_ok());
305 let max_name = "a".repeat(MAX_SAVEPOINT_NAME_LEN);
307 assert!(validate_savepoint_name(&max_name).is_ok());
308 }
309
310 #[test]
311 fn test_validate_savepoint_name_rejects_invalid_names() {
312 assert!(validate_savepoint_name("sp1; DROP TABLE").is_err());
313 assert!(validate_savepoint_name("my savepoint").is_err());
314 assert!(validate_savepoint_name("\"quoted\"").is_err());
315 assert!(validate_savepoint_name("").is_err());
316 assert!(validate_savepoint_name("1leading_digit").is_err());
317 assert!(validate_savepoint_name("has-dash").is_err());
318 assert!(validate_savepoint_name("has.dot").is_err());
319 let too_long = "a".repeat(MAX_SAVEPOINT_NAME_LEN + 1);
321 assert!(validate_savepoint_name(&too_long).is_err());
322 }
323}