systemprompt_database/services/postgres/
mod.rs1pub 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}