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