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 introspection;
15pub mod transaction;
16
17use async_trait::async_trait;
18use sqlx::Executor;
19use sqlx::postgres::PgPool;
20use std::sync::Arc;
21
22use super::provider::DatabaseProvider;
23use crate::error::DatabaseResult;
24use crate::models::{
25    DatabaseInfo, DatabaseTransaction, JsonRow, QueryResult, QuerySelector, ToDbValue,
26};
27use conversion::{bind_params, row_to_json, rows_to_result};
28use transaction::PostgresTransaction;
29
30#[derive(Debug)]
31pub struct PostgresProvider {
32    pool: Arc<PgPool>,
33}
34
35impl PostgresProvider {
36    pub async fn new(database_url: &str) -> DatabaseResult<Self> {
37        Self::new_with_pool(database_url, &connection::PoolConfig::default()).await
38    }
39
40    pub async fn new_with_pool(
41        database_url: &str,
42        pool_config: &connection::PoolConfig,
43    ) -> DatabaseResult<Self> {
44        let connect_options = connection::connect_options(database_url)?;
45
46        let pool = connection::connect_with_retry(
47            connection::build_pool_options(pool_config),
48            connect_options,
49        )
50        .await?;
51
52        Ok(Self {
53            pool: Arc::new(pool),
54        })
55    }
56
57    #[must_use]
58    pub const fn from_pool(pool: Arc<PgPool>) -> Self {
59        Self { pool }
60    }
61
62    #[must_use]
63    pub fn pool(&self) -> &PgPool {
64        &self.pool
65    }
66}
67
68#[async_trait]
69impl DatabaseProvider for PostgresProvider {
70    fn get_postgres_pool(&self) -> Arc<PgPool> {
71        Arc::clone(&self.pool)
72    }
73
74    async fn execute(
75        &self,
76        query: &dyn QuerySelector,
77        params: &[&dyn ToDbValue],
78    ) -> DatabaseResult<u64> {
79        let sql = query.select_query();
80        let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
81        let query_obj = bind_params(query_obj, params);
82
83        let result = query_obj.execute(&*self.pool).await?;
84
85        Ok(result.rows_affected())
86    }
87
88    async fn execute_raw(&self, sql: &str) -> DatabaseResult<()> {
89        let mut conn = self.pool.acquire().await?;
90
91        conn.execute(sqlx::AssertSqlSafe(sql.to_owned())).await?;
92
93        Ok(())
94    }
95
96    async fn fetch_all(
97        &self,
98        query: &dyn QuerySelector,
99        params: &[&dyn ToDbValue],
100    ) -> DatabaseResult<Vec<JsonRow>> {
101        let sql = query.select_query();
102        let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
103        let query_obj = bind_params(query_obj, params);
104
105        let rows = query_obj.fetch_all(&*self.pool).await?;
106
107        Ok(rows.iter().map(row_to_json).collect())
108    }
109
110    async fn fetch_one(
111        &self,
112        query: &dyn QuerySelector,
113        params: &[&dyn ToDbValue],
114    ) -> DatabaseResult<JsonRow> {
115        let sql = query.select_query();
116        let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
117        let query_obj = bind_params(query_obj, params);
118
119        let row = query_obj.fetch_one(&*self.pool).await?;
120
121        Ok(row_to_json(&row))
122    }
123
124    async fn fetch_optional(
125        &self,
126        query: &dyn QuerySelector,
127        params: &[&dyn ToDbValue],
128    ) -> DatabaseResult<Option<JsonRow>> {
129        let sql = query.select_query();
130        let query_obj = sqlx::query(sqlx::AssertSqlSafe(sql));
131        let query_obj = bind_params(query_obj, params);
132
133        let row = query_obj.fetch_optional(&*self.pool).await?;
134
135        Ok(row.map(|r| row_to_json(&r)))
136    }
137
138    async fn begin_transaction(&self) -> DatabaseResult<Box<dyn DatabaseTransaction>> {
139        let tx = self.pool.begin().await?;
140
141        Ok(Box::new(PostgresTransaction::new(tx)))
142    }
143
144    async fn get_database_info(&self) -> DatabaseResult<DatabaseInfo> {
145        introspection::get_database_info(&self.pool).await
146    }
147
148    async fn test_connection(&self) -> DatabaseResult<()> {
149        sqlx::query("SELECT 1").fetch_one(&*self.pool).await?;
150        Ok(())
151    }
152
153    async fn execute_batch(&self, sql: &str) -> DatabaseResult<()> {
154        let statements = crate::services::SqlExecutor::parse_sql_statements(sql)?;
155        for statement in statements {
156            sqlx::query(sqlx::AssertSqlSafe(statement))
157                .execute(&*self.pool)
158                .await?;
159        }
160        Ok(())
161    }
162
163    async fn query_raw(&self, query: &dyn QuerySelector) -> DatabaseResult<QueryResult> {
164        let sql = query.select_query();
165        let start = std::time::Instant::now();
166
167        let rows = sqlx::query(sqlx::AssertSqlSafe(sql))
168            .fetch_all(&*self.pool)
169            .await?;
170
171        Ok(rows_to_result(rows, start))
172    }
173
174    async fn query_raw_with(
175        &self,
176        query: &dyn QuerySelector,
177        params: &[&dyn ToDbValue],
178    ) -> DatabaseResult<QueryResult> {
179        let sql = query.select_query();
180        let start = std::time::Instant::now();
181
182        let query_obj = bind_params(sqlx::query(sqlx::AssertSqlSafe(sql)), params);
183        let rows = query_obj.fetch_all(&*self.pool).await?;
184
185        Ok(rows_to_result(rows, start))
186    }
187}