systemprompt_database/services/postgres/
mod.rs1pub mod connection;
13pub mod conversion;
14mod ext;
15mod introspection;
16pub mod transaction;
17
18use async_trait::async_trait;
19use sqlx::Executor;
20use sqlx::postgres::{PgConnectOptions, PgPool, PgSslMode};
21use std::str::FromStr;
22use std::sync::Arc;
23
24use super::provider::DatabaseProvider;
25use crate::error::{DatabaseResult, RepositoryError};
26use crate::models::{
27 DatabaseInfo, DatabaseTransaction, DbValue, JsonRow, QueryResult, QuerySelector, ToDbValue,
28};
29use conversion::{bind_params, row_to_json, rows_to_result};
30use transaction::PostgresTransaction;
31
32#[derive(Debug)]
33pub struct PostgresProvider {
34 pool: Arc<PgPool>,
35}
36
37impl PostgresProvider {
38 pub async fn new(database_url: &str) -> DatabaseResult<Self> {
39 Self::new_with_pool(database_url, &connection::PoolConfig::default()).await
40 }
41
42 pub async fn new_with_pool(
43 database_url: &str,
44 pool_config: &connection::PoolConfig,
45 ) -> DatabaseResult<Self> {
46 let mut connect_options = PgConnectOptions::from_str(database_url)?;
47
48 let ssl_mode = if database_url.contains("sslmode=require") {
49 PgSslMode::Require
50 } else if database_url.contains("sslmode=disable") {
51 PgSslMode::Disable
52 } else {
53 PgSslMode::Prefer
54 };
55
56 connect_options = connect_options
57 .application_name("systemprompt")
58 .statement_cache_capacity(0)
59 .ssl_mode(ssl_mode)
60 .options([("client_min_messages", "warning")]);
61
62 let pool = connection::connect_with_retry(
63 connection::build_pool_options(pool_config),
64 connect_options,
65 )
66 .await?;
67
68 Ok(Self {
69 pool: Arc::new(pool),
70 })
71 }
72
73 #[must_use]
74 pub const fn from_pool(pool: Arc<PgPool>) -> Self {
75 Self { pool }
76 }
77
78 #[must_use]
79 pub fn pool(&self) -> &PgPool {
80 &self.pool
81 }
82}
83
84#[async_trait]
85impl DatabaseProvider for PostgresProvider {
86 fn get_postgres_pool(&self) -> Option<Arc<PgPool>> {
87 Some(Arc::clone(&self.pool))
88 }
89
90 async fn execute(
91 &self,
92 query: &dyn QuerySelector,
93 params: &[&dyn ToDbValue],
94 ) -> DatabaseResult<u64> {
95 let sql = query.select_query();
96 let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
97 let query_obj = bind_params(query_obj, params);
98
99 let result = query_obj.execute(&*self.pool).await?;
100
101 Ok(result.rows_affected())
102 }
103
104 async fn execute_raw(&self, sql: &str) -> DatabaseResult<()> {
105 let mut conn = self.pool.acquire().await?;
106
107 conn.execute(sqlx::AssertSqlSafe(sql.to_owned())).await?;
108
109 Ok(())
110 }
111
112 async fn fetch_all(
113 &self,
114 query: &dyn QuerySelector,
115 params: &[&dyn ToDbValue],
116 ) -> DatabaseResult<Vec<JsonRow>> {
117 let sql = query.select_query();
118 let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
119 let query_obj = bind_params(query_obj, params);
120
121 let rows = query_obj.fetch_all(&*self.pool).await?;
122
123 Ok(rows.iter().map(row_to_json).collect())
124 }
125
126 async fn fetch_one(
127 &self,
128 query: &dyn QuerySelector,
129 params: &[&dyn ToDbValue],
130 ) -> DatabaseResult<JsonRow> {
131 let sql = query.select_query();
132 let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
133 let query_obj = bind_params(query_obj, params);
134
135 let row = query_obj.fetch_one(&*self.pool).await?;
136
137 Ok(row_to_json(&row))
138 }
139
140 async fn fetch_optional(
141 &self,
142 query: &dyn QuerySelector,
143 params: &[&dyn ToDbValue],
144 ) -> DatabaseResult<Option<JsonRow>> {
145 let sql = query.select_query();
146 let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
147 let query_obj = bind_params(query_obj, params);
148
149 let row = query_obj.fetch_optional(&*self.pool).await?;
150
151 Ok(row.map(|r| row_to_json(&r)))
152 }
153
154 async fn fetch_scalar_value(
155 &self,
156 query: &dyn QuerySelector,
157 params: &[&dyn ToDbValue],
158 ) -> DatabaseResult<DbValue> {
159 let row = self.fetch_one(query, params).await?;
160
161 let first_value = row
162 .values()
163 .next()
164 .ok_or_else(|| RepositoryError::invalid_state("No columns in result"))?;
165
166 let db_value = match first_value {
167 serde_json::Value::String(s) => DbValue::String(s.clone()),
168 serde_json::Value::Number(n) => n
169 .as_i64()
170 .map(DbValue::Int)
171 .or_else(|| n.as_f64().map(DbValue::Float))
172 .unwrap_or(DbValue::NullFloat),
173 serde_json::Value::Bool(b) => DbValue::Bool(*b),
174 serde_json::Value::Null => DbValue::NullString,
175 serde_json::Value::Array(_) | serde_json::Value::Object(_) => {
176 return Err(RepositoryError::invalid_state("Unsupported value type"));
177 },
178 };
179
180 Ok(db_value)
181 }
182
183 async fn begin_transaction(&self) -> DatabaseResult<Box<dyn DatabaseTransaction>> {
184 let tx = self.pool.begin().await?;
185
186 Ok(Box::new(PostgresTransaction::new(tx)))
187 }
188
189 async fn get_database_info(&self) -> DatabaseResult<DatabaseInfo> {
190 introspection::get_database_info(&self.pool).await
191 }
192
193 async fn test_connection(&self) -> DatabaseResult<()> {
194 sqlx::query("SELECT 1").fetch_one(&*self.pool).await?;
195 Ok(())
196 }
197
198 async fn execute_batch(&self, sql: &str) -> DatabaseResult<()> {
199 let statements = crate::services::SqlExecutor::parse_sql_statements(sql)?;
200 for statement in statements {
201 sqlx::query(sqlx::AssertSqlSafe(statement))
202 .execute(&*self.pool)
203 .await?;
204 }
205 Ok(())
206 }
207
208 async fn query_raw(&self, query: &dyn QuerySelector) -> DatabaseResult<QueryResult> {
209 let sql = query.select_query();
210 let start = std::time::Instant::now();
211
212 let rows = sqlx::query(sqlx::AssertSqlSafe(sql))
213 .fetch_all(&*self.pool)
214 .await?;
215
216 Ok(rows_to_result(rows, start))
217 }
218
219 async fn query_raw_with(
220 &self,
221 query: &dyn QuerySelector,
222 params: &[&dyn ToDbValue],
223 ) -> DatabaseResult<QueryResult> {
224 let sql = query.select_query();
225 let start = std::time::Instant::now();
226
227 let query_obj = bind_params(sqlx::query(sqlx::AssertSqlSafe(sql)), params);
228 let rows = query_obj.fetch_all(&*self.pool).await?;
229
230 Ok(rows_to_result(rows, start))
231 }
232}