Skip to main content

systemprompt_database/services/postgres/
mod.rs

1//! `PostgreSQL` implementation of [`crate::services::DatabaseProvider`].
2//!
3//! This module is part of the documented sqlx allowlist: every `sqlx::query(_)`
4//! call here either binds a [`crate::models::QuerySelector`] string supplied
5//! at runtime (extension-defined SQL, dynamic admin queries) or executes
6//! `SELECT 1` for connection probing. Static SQL goes through the verified
7//! macros elsewhere.
8//!
9//! Copyright (c) systemprompt.io — Business Source License 1.1.
10//! See <https://systemprompt.io> for licensing details.
11
12pub 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}