Skip to main content

prax_postgres/
connection.rs

1//! PostgreSQL connection wrapper.
2
3use 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
14/// A wrapper around a PostgreSQL connection with statement caching.
15pub struct PgConnection {
16    client: Object,
17    statement_cache: Arc<PreparedStatementCache>,
18}
19
20impl PgConnection {
21    /// Create a new connection wrapper.
22    pub(crate) fn new(client: Object, statement_cache: Arc<PreparedStatementCache>) -> Self {
23        Self {
24            client,
25            statement_cache,
26        }
27    }
28
29    /// Execute a query and return all rows.
30    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        // Try to get a cached prepared statement
38        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    /// Execute a query and return exactly one row.
48    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    /// Execute a query and return zero or one row.
65    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    /// Execute a statement and return the number of affected rows.
82    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    /// Execute a batch of statements in a single round-trip.
99    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    /// Begin a transaction.
106    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    /// Get the underlying tokio-postgres client.
116    ///
117    /// This is useful for advanced operations not covered by this wrapper.
118    pub fn inner(&self) -> &Object {
119        &self.client
120    }
121
122    /// Execute a query using the prepared statement cache.
123    ///
124    /// This is an alias for `query` that makes it explicit that statement caching
125    /// is being used. All query methods already use prepared statement caching,
126    /// but this method name makes it more explicit for benchmark comparisons.
127    #[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    /// Execute a raw query without using the prepared statement cache.
137    ///
138    /// This is useful for one-off queries where the overhead of preparing
139    /// a statement isn't worth it.
140    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    /// Execute a raw query and return zero or one row without using statement cache.
151    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
162/// Maximum allowed savepoint name length (matches PostgreSQL's `NAMEDATALEN - 1`).
163const MAX_SAVEPOINT_NAME_LEN: usize = 63;
164
165/// Validate a savepoint name before it is interpolated into SQL.
166///
167/// Savepoint identifiers cannot be parameterized, so they must match the
168/// whitelist pattern `^[A-Za-z_][A-Za-z0-9_]*$` to prevent SQL injection.
169fn 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
177/// A PostgreSQL transaction.
178pub struct PgTransaction<'a> {
179    txn: deadpool_postgres::Transaction<'a>,
180    statement_cache: Arc<PreparedStatementCache>,
181}
182
183impl<'a> PgTransaction<'a> {
184    /// Execute a query and return all rows.
185    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    /// Execute a query and return exactly one row.
202    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    /// Execute a query and return zero or one row.
217    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    /// Execute a statement and return the number of affected rows.
232    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    /// Create a savepoint.
247    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    /// Rollback to a savepoint.
257    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    /// Release a savepoint.
267    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    /// Commit the transaction.
277    pub async fn commit(self) -> PgResult<()> {
278        debug!("Committing transaction");
279        self.txn.commit().await?;
280        Ok(())
281    }
282
283    /// Rollback the transaction.
284    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    // Integration tests would require a real PostgreSQL connection
296    // Unit tests for connection wrapper are limited without mocking
297
298    #[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        // 63 chars (the max) is accepted
306        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        // 64 chars exceeds the limit
320        let too_long = "a".repeat(MAX_SAVEPOINT_NAME_LEN + 1);
321        assert!(validate_savepoint_name(&too_long).is_err());
322    }
323}