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::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}