1use crate::config::DatabaseBackend;
4use crate::error::SqlxResult;
5use crate::types::quote_identifier;
6use sqlx::Row;
7use sqlx::postgres::{PgPool, PgRow};
8
9pub struct PgHelpers;
11
12impl PgHelpers {
13 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 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 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 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 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 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 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 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
136pub struct AdvisoryLock;
138
139impl AdvisoryLock {
140 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 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 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 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 assert_eq!(
191 PgHelpers::json_path("data", &["na'me"]),
192 "\"data\"->>'na''me'"
193 );
194 }
195}