Skip to main content

systemprompt_database/services/
database.rs

1//! Top-level [`Database`] handle that owns one or two
2//! [`DatabaseProvider`] instances (read + optional write) and exposes the
3//! query and transaction surface.
4//!
5//! Copyright (c) systemprompt.io — Business Source License 1.1.
6//! See <https://systemprompt.io> for licensing details.
7
8use super::postgres::PostgresProvider;
9use super::postgres::connection::PoolConfig;
10use super::provider::DatabaseProvider;
11use crate::error::DatabaseResult;
12use crate::models::{DatabaseInfo, QueryResult};
13use std::sync::Arc;
14
15pub struct Database {
16    provider: Arc<dyn DatabaseProvider>,
17    write_provider: Option<Arc<dyn DatabaseProvider>>,
18}
19
20impl std::fmt::Debug for Database {
21    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
22        f.debug_struct("Database")
23            .field("backend", &"PostgreSQL")
24            .finish()
25    }
26}
27
28impl Database {
29    pub async fn new_postgres(url: &str) -> DatabaseResult<Self> {
30        let provider = PostgresProvider::new(url).await?;
31        Ok(Self {
32            provider: Arc::new(provider),
33            write_provider: None,
34        })
35    }
36
37    pub async fn connect(
38        read_url: &str,
39        write_url: Option<&str>,
40        pool: &PoolConfig,
41    ) -> DatabaseResult<Self> {
42        let provider: Arc<dyn DatabaseProvider> =
43            Arc::new(PostgresProvider::new_with_pool(read_url, pool).await?);
44
45        let write_provider: Option<Arc<dyn DatabaseProvider>> = match write_url {
46            Some(url) => Some(Arc::new(PostgresProvider::new_with_pool(url, pool).await?)),
47            None => None,
48        };
49
50        Ok(Self {
51            provider,
52            write_provider,
53        })
54    }
55
56    #[must_use]
57    pub fn from_pools(read: Arc<sqlx::PgPool>, write: Option<Arc<sqlx::PgPool>>) -> Self {
58        let write_provider = write.map(|pool| -> Arc<dyn DatabaseProvider> {
59            Arc::new(PostgresProvider::from_pool(pool))
60        });
61        Self {
62            provider: Arc::new(PostgresProvider::from_pool(read)),
63            write_provider,
64        }
65    }
66
67    #[must_use]
68    pub fn read(&self) -> &dyn DatabaseProvider {
69        self.provider.as_ref()
70    }
71
72    #[must_use]
73    pub fn write(&self) -> &dyn DatabaseProvider {
74        self.write_provider
75            .as_deref()
76            .unwrap_or_else(|| self.provider.as_ref())
77    }
78
79    #[must_use]
80    pub fn pool(&self) -> Arc<sqlx::PgPool> {
81        self.read().get_postgres_pool()
82    }
83
84    #[expect(
85        clippy::unnecessary_wraps,
86        reason = "every layer threads `?` through this accessor; collapsing its callers onto \
87                  `pool()` is a workspace-wide mechanical change scheduled after 0.53.0"
88    )]
89    pub fn pool_arc(&self) -> DatabaseResult<Arc<sqlx::PgPool>> {
90        Ok(self.pool())
91    }
92
93    #[must_use]
94    pub fn write_pool(&self) -> Arc<sqlx::PgPool> {
95        self.write().get_postgres_pool()
96    }
97
98    #[expect(
99        clippy::unnecessary_wraps,
100        reason = "every layer threads `?` through this accessor; collapsing its callers onto \
101                  `write_pool()` is a workspace-wide mechanical change scheduled after 0.53.0"
102    )]
103    pub fn write_pool_arc(&self) -> DatabaseResult<Arc<sqlx::PgPool>> {
104        Ok(self.write_pool())
105    }
106
107    #[must_use]
108    pub fn has_write_pool(&self) -> bool {
109        self.write_provider.is_some()
110    }
111
112    pub async fn execute_batch(&self, sql: &str) -> DatabaseResult<()> {
113        self.write().execute_batch(sql).await
114    }
115
116    pub async fn get_info(&self) -> DatabaseResult<DatabaseInfo> {
117        self.read().get_database_info().await
118    }
119
120    pub async fn test_connection(&self) -> DatabaseResult<()> {
121        self.provider.test_connection().await?;
122        if let Some(wp) = &self.write_provider {
123            wp.test_connection().await?;
124        }
125        Ok(())
126    }
127
128    pub async fn begin(&self) -> DatabaseResult<sqlx::Transaction<'_, sqlx::Postgres>> {
129        self.write_pool().begin().await.map_err(Into::into)
130    }
131
132    pub async fn begin_scoped(
133        &self,
134        scope: &systemprompt_models::RequestScope,
135    ) -> DatabaseResult<sqlx::Transaction<'static, sqlx::Postgres>> {
136        super::scoped_transaction::begin_scoped(&self.write_pool(), scope).await
137    }
138}
139
140pub type DbPool = Arc<Database>;
141
142pub trait DatabaseExt {
143    fn database(&self) -> Arc<Database>;
144}
145
146impl DatabaseExt for Arc<Database> {
147    fn database(&self) -> Arc<Database> {
148        Self::clone(self)
149    }
150}
151
152#[async_trait::async_trait]
153impl DatabaseProvider for Database {
154    fn get_postgres_pool(&self) -> Arc<sqlx::PgPool> {
155        self.read().get_postgres_pool()
156    }
157
158    async fn execute(
159        &self,
160        query: &dyn crate::models::QuerySelector,
161        params: &[&dyn crate::models::ToDbValue],
162    ) -> DatabaseResult<u64> {
163        self.write().execute(query, params).await
164    }
165
166    async fn execute_raw(&self, sql: &str) -> DatabaseResult<()> {
167        self.write().execute_raw(sql).await
168    }
169
170    async fn fetch_all(
171        &self,
172        query: &dyn crate::models::QuerySelector,
173        params: &[&dyn crate::models::ToDbValue],
174    ) -> DatabaseResult<Vec<crate::models::JsonRow>> {
175        self.read().fetch_all(query, params).await
176    }
177
178    async fn fetch_one(
179        &self,
180        query: &dyn crate::models::QuerySelector,
181        params: &[&dyn crate::models::ToDbValue],
182    ) -> DatabaseResult<crate::models::JsonRow> {
183        self.read().fetch_one(query, params).await
184    }
185
186    async fn fetch_optional(
187        &self,
188        query: &dyn crate::models::QuerySelector,
189        params: &[&dyn crate::models::ToDbValue],
190    ) -> DatabaseResult<Option<crate::models::JsonRow>> {
191        self.read().fetch_optional(query, params).await
192    }
193
194    async fn begin_transaction(
195        &self,
196    ) -> DatabaseResult<Box<dyn crate::models::DatabaseTransaction>> {
197        self.write().begin_transaction().await
198    }
199
200    async fn get_database_info(&self) -> DatabaseResult<DatabaseInfo> {
201        self.read().get_database_info().await
202    }
203
204    async fn test_connection(&self) -> DatabaseResult<()> {
205        self.read().test_connection().await
206    }
207
208    async fn execute_batch(&self, sql: &str) -> DatabaseResult<()> {
209        self.write().execute_batch(sql).await
210    }
211
212    async fn query_raw(
213        &self,
214        query: &dyn crate::models::QuerySelector,
215    ) -> DatabaseResult<QueryResult> {
216        self.read().query_raw(query).await
217    }
218
219    async fn query_raw_with(
220        &self,
221        query: &dyn crate::models::QuerySelector,
222        params: &[&dyn crate::models::ToDbValue],
223    ) -> DatabaseResult<QueryResult> {
224        self.read().query_raw_with(query, params).await
225    }
226}