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    #[must_use]
85    pub fn write_pool(&self) -> Arc<sqlx::PgPool> {
86        self.write().get_postgres_pool()
87    }
88
89    #[must_use]
90    pub fn has_write_pool(&self) -> bool {
91        self.write_provider.is_some()
92    }
93
94    pub async fn execute_batch(&self, sql: &str) -> DatabaseResult<()> {
95        self.write().execute_batch(sql).await
96    }
97
98    pub async fn get_info(&self) -> DatabaseResult<DatabaseInfo> {
99        self.read().get_database_info().await
100    }
101
102    pub async fn test_connection(&self) -> DatabaseResult<()> {
103        self.provider.test_connection().await?;
104        if let Some(wp) = &self.write_provider {
105            wp.test_connection().await?;
106        }
107        Ok(())
108    }
109
110    pub async fn begin(&self) -> DatabaseResult<sqlx::Transaction<'_, sqlx::Postgres>> {
111        self.write_pool().begin().await.map_err(Into::into)
112    }
113}
114
115pub type DbPool = Arc<Database>;
116
117pub trait DatabaseExt {
118    fn database(&self) -> Arc<Database>;
119}
120
121impl DatabaseExt for Arc<Database> {
122    fn database(&self) -> Arc<Database> {
123        Self::clone(self)
124    }
125}
126
127#[async_trait::async_trait]
128impl DatabaseProvider for Database {
129    fn get_postgres_pool(&self) -> Arc<sqlx::PgPool> {
130        self.read().get_postgres_pool()
131    }
132
133    async fn execute(
134        &self,
135        query: &dyn crate::models::QuerySelector,
136        params: &[&dyn crate::models::ToDbValue],
137    ) -> DatabaseResult<u64> {
138        self.write().execute(query, params).await
139    }
140
141    async fn execute_raw(&self, sql: &str) -> DatabaseResult<()> {
142        self.write().execute_raw(sql).await
143    }
144
145    async fn fetch_all(
146        &self,
147        query: &dyn crate::models::QuerySelector,
148        params: &[&dyn crate::models::ToDbValue],
149    ) -> DatabaseResult<Vec<crate::models::JsonRow>> {
150        self.read().fetch_all(query, params).await
151    }
152
153    async fn fetch_one(
154        &self,
155        query: &dyn crate::models::QuerySelector,
156        params: &[&dyn crate::models::ToDbValue],
157    ) -> DatabaseResult<crate::models::JsonRow> {
158        self.read().fetch_one(query, params).await
159    }
160
161    async fn fetch_optional(
162        &self,
163        query: &dyn crate::models::QuerySelector,
164        params: &[&dyn crate::models::ToDbValue],
165    ) -> DatabaseResult<Option<crate::models::JsonRow>> {
166        self.read().fetch_optional(query, params).await
167    }
168
169    async fn begin_transaction(
170        &self,
171    ) -> DatabaseResult<Box<dyn crate::models::DatabaseTransaction>> {
172        self.write().begin_transaction().await
173    }
174
175    async fn get_database_info(&self) -> DatabaseResult<DatabaseInfo> {
176        self.read().get_database_info().await
177    }
178
179    async fn test_connection(&self) -> DatabaseResult<()> {
180        self.read().test_connection().await
181    }
182
183    async fn execute_batch(&self, sql: &str) -> DatabaseResult<()> {
184        self.write().execute_batch(sql).await
185    }
186
187    async fn query_raw(
188        &self,
189        query: &dyn crate::models::QuerySelector,
190    ) -> DatabaseResult<QueryResult> {
191        self.read().query_raw(query).await
192    }
193
194    async fn query_raw_with(
195        &self,
196        query: &dyn crate::models::QuerySelector,
197        params: &[&dyn crate::models::ToDbValue],
198    ) -> DatabaseResult<QueryResult> {
199        self.read().query_raw_with(query, params).await
200    }
201}