Skip to main content

prax_sqlx/
postgres.rs

1//! PostgreSQL-specific functionality for SQLx.
2
3use crate::config::DatabaseBackend;
4use crate::error::SqlxResult;
5use crate::types::quote_identifier;
6use sqlx::Row;
7use sqlx::postgres::{PgPool, PgRow};
8
9/// PostgreSQL-specific query helpers.
10pub struct PgHelpers;
11
12impl PgHelpers {
13    /// Execute a query with RETURNING clause.
14    pub async fn query_returning(pool: &PgPool, sql: &str) -> SqlxResult<Vec<PgRow>> {
15        let rows = sqlx::query(sql).fetch_all(pool).await?;
16        Ok(rows)
17    }
18
19    /// Build INSERT ... ON CONFLICT SQL (does not execute; run it via your pool).
20    ///
21    /// The `pool` parameter is currently unused; it is reserved for a future
22    /// executing variant. Identifiers are quoted via `quote_identifier`;
23    /// pass trusted identifiers only, not arbitrary user input.
24    pub async fn upsert(
25        _pool: &PgPool,
26        table: &str,
27        columns: &[&str],
28        conflict_columns: &[&str],
29        update_columns: &[&str],
30    ) -> SqlxResult<String> {
31        let table = quote_identifier(DatabaseBackend::Postgres, table);
32        let cols = columns
33            .iter()
34            .map(|c| quote_identifier(DatabaseBackend::Postgres, c))
35            .collect::<Vec<_>>()
36            .join(", ");
37        let placeholders: Vec<String> = (1..=columns.len()).map(|i| format!("${}", i)).collect();
38        let vals = placeholders.join(", ");
39        let conflict = conflict_columns
40            .iter()
41            .map(|c| quote_identifier(DatabaseBackend::Postgres, c))
42            .collect::<Vec<_>>()
43            .join(", ");
44        let updates: Vec<String> = update_columns
45            .iter()
46            .map(|c| {
47                let col = quote_identifier(DatabaseBackend::Postgres, c);
48                format!("{} = EXCLUDED.{}", col, col)
49            })
50            .collect();
51        let update_clause = updates.join(", ");
52
53        let sql = format!(
54            "INSERT INTO {} ({}) VALUES ({}) ON CONFLICT ({}) DO UPDATE SET {} RETURNING *",
55            table, cols, vals, conflict, update_clause
56        );
57
58        Ok(sql)
59    }
60
61    /// Generate a PostgreSQL array literal.
62    ///
63    /// Embedded single quotes in values are escaped (`'` -> `''`). Prefer
64    /// bound parameters for untrusted values.
65    pub fn array_literal<T: std::fmt::Display>(values: &[T]) -> String {
66        let items: Vec<String> = values
67            .iter()
68            .map(|v| format!("'{}'", v.to_string().replace('\'', "''")))
69            .collect();
70        format!("ARRAY[{}]", items.join(", "))
71    }
72
73    /// Generate a PostgreSQL JSON/JSONB path expression.
74    ///
75    /// The column is quoted via `quote_identifier` and embedded single quotes
76    /// in path elements are escaped (`'` -> `''`). Pass trusted identifiers
77    /// only, not arbitrary user input.
78    pub fn json_path(column: &str, path: &[&str]) -> String {
79        let column = quote_identifier(DatabaseBackend::Postgres, column);
80        if path.is_empty() {
81            column
82        } else {
83            let path_str: Vec<String> = path
84                .iter()
85                .map(|p| format!("'{}'", p.replace('\'', "''")))
86                .collect();
87            format!("{}->>{}", column, path_str.join("->"))
88        }
89    }
90
91    /// Check if a PostgreSQL extension is available.
92    pub async fn has_extension(pool: &PgPool, extension: &str) -> SqlxResult<bool> {
93        let sql = "SELECT EXISTS(SELECT 1 FROM pg_extension WHERE extname = $1)";
94        let row = sqlx::query(sql).bind(extension).fetch_one(pool).await?;
95        let exists: bool = row.try_get(0)?;
96        Ok(exists)
97    }
98
99    /// Get PostgreSQL version.
100    pub async fn version(pool: &PgPool) -> SqlxResult<String> {
101        let sql = "SELECT version()";
102        let row = sqlx::query(sql).fetch_one(pool).await?;
103        let version: String = row.try_get(0)?;
104        Ok(version)
105    }
106
107    /// Execute LISTEN for notifications.
108    ///
109    /// The channel is a PostgreSQL identifier and is quoted via
110    /// `quote_identifier`; pass trusted channel names only.
111    pub async fn listen(pool: &PgPool, channel: &str) -> SqlxResult<()> {
112        let sql = format!(
113            "LISTEN {}",
114            quote_identifier(DatabaseBackend::Postgres, channel)
115        );
116        sqlx::query(&sql).execute(pool).await?;
117        Ok(())
118    }
119
120    /// Execute NOTIFY.
121    ///
122    /// The channel is a PostgreSQL identifier and is quoted via
123    /// `quote_identifier`; single quotes in the payload are escaped
124    /// (`'` -> `''`). Pass trusted channel names only.
125    pub async fn notify(pool: &PgPool, channel: &str, payload: &str) -> SqlxResult<()> {
126        let sql = format!(
127            "NOTIFY {}, '{}'",
128            quote_identifier(DatabaseBackend::Postgres, channel),
129            payload.replace('\'', "''")
130        );
131        sqlx::query(&sql).execute(pool).await?;
132        Ok(())
133    }
134}
135
136/// PostgreSQL advisory lock helpers.
137pub struct AdvisoryLock;
138
139impl AdvisoryLock {
140    /// Acquire an advisory lock (blocking).
141    pub async fn acquire(pool: &PgPool, key: i64) -> SqlxResult<()> {
142        sqlx::query("SELECT pg_advisory_lock($1)")
143            .bind(key)
144            .execute(pool)
145            .await?;
146        Ok(())
147    }
148
149    /// Try to acquire an advisory lock (non-blocking).
150    pub async fn try_acquire(pool: &PgPool, key: i64) -> SqlxResult<bool> {
151        let row = sqlx::query("SELECT pg_try_advisory_lock($1)")
152            .bind(key)
153            .fetch_one(pool)
154            .await?;
155        let acquired: bool = row.try_get(0)?;
156        Ok(acquired)
157    }
158
159    /// Release an advisory lock.
160    pub async fn release(pool: &PgPool, key: i64) -> SqlxResult<()> {
161        sqlx::query("SELECT pg_advisory_unlock($1)")
162            .bind(key)
163            .execute(pool)
164            .await?;
165        Ok(())
166    }
167}
168
169#[cfg(test)]
170mod tests {
171    use super::*;
172
173    #[test]
174    fn test_array_literal() {
175        assert_eq!(PgHelpers::array_literal(&[1, 2, 3]), "ARRAY['1', '2', '3']");
176        assert_eq!(PgHelpers::array_literal(&["a", "b"]), "ARRAY['a', 'b']");
177        // Embedded single quotes are escaped.
178        assert_eq!(PgHelpers::array_literal(&["a'b"]), "ARRAY['a''b']");
179    }
180
181    #[test]
182    fn test_json_path() {
183        assert_eq!(PgHelpers::json_path("data", &[]), "\"data\"");
184        assert_eq!(PgHelpers::json_path("data", &["name"]), "\"data\"->>'name'");
185        assert_eq!(
186            PgHelpers::json_path("data", &["user", "name"]),
187            "\"data\"->>'user'->'name'"
188        );
189        // Embedded single quotes in path elements are escaped.
190        assert_eq!(
191            PgHelpers::json_path("data", &["na'me"]),
192            "\"data\"->>'na''me'"
193        );
194    }
195}