1mod conv;
2mod errors;
3mod manager;
4pub mod migration;
5pub mod sqlite;
6
7use std::num::NonZeroUsize;
8use std::path::Path;
9
10pub use conv::{DatabaseTypeConversionError, SqlTypeConvert};
11use diesel::{RunQueryDsl, SqliteConnection};
12pub use errors::{DatabaseError, SchemaVerificationError};
13pub use manager::{ConnectionManager, ConnectionManagerError, configure_connection_on_creation};
14use miden_node_tracing::Instrument;
15
16pub type Result<T, E = DatabaseError> = std::result::Result<T, E>;
17
18pub fn default_connection_pool_size() -> NonZeroUsize {
23 let available_cores = std::thread::available_parallelism().map_or(1, NonZeroUsize::get);
24 let connection_count = available_cores.saturating_mul(2);
25 NonZeroUsize::new(connection_count).expect("connection count must be non-zero")
26}
27
28#[derive(Clone)]
31pub struct Db {
32 pool: deadpool_diesel::Pool<ConnectionManager, deadpool::managed::Object<ConnectionManager>>,
33}
34
35impl Db {
36 pub fn new(database_filepath: &Path) -> Result<Self, DatabaseError> {
38 Self::new_with_pool_size(database_filepath, default_connection_pool_size())
39 }
40
41 pub fn new_with_pool_size(
43 database_filepath: &Path,
44 connection_pool_size: NonZeroUsize,
45 ) -> Result<Self, DatabaseError> {
46 let manager = ConnectionManager::new(database_filepath.to_str().unwrap());
47 let pool = deadpool_diesel::Pool::builder(manager)
48 .max_size(connection_pool_size.get())
49 .build()?;
50 Ok(Self { pool })
51 }
52
53 pub async fn pinned_connection(&self) -> Result<PinnedConnection, DatabaseError> {
59 let conn = self
60 .pool
61 .get()
62 .in_current_span()
63 .await
64 .map_err(|e| DatabaseError::ConnectionPoolObtainError(Box::new(e)))?;
65 Ok(PinnedConnection { conn })
66 }
67
68 pub async fn transact<R, E, Q, M>(&self, msg: M, query: Q) -> std::result::Result<R, E>
70 where
71 Q: Send + for<'a> FnOnce(&'a mut SqliteConnection) -> std::result::Result<R, E> + 'static,
72 R: Send + 'static,
73 M: Send + ToString,
74 E: From<diesel::result::Error>,
75 E: From<DatabaseError>,
76 E: std::error::Error + Send + Sync + 'static,
77 {
78 self.pinned_connection().await.map_err(E::from)?.transact(msg, query).await
79 }
80
81 pub async fn query<R, E, Q, M>(&self, msg: M, query: Q) -> std::result::Result<R, E>
83 where
84 Q: Send + FnOnce(&mut SqliteConnection) -> std::result::Result<R, E> + 'static,
85 R: Send + 'static,
86 M: Send + ToString,
87 E: From<DatabaseError>,
88 E: std::error::Error + Send + Sync + 'static,
89 {
90 self.pinned_connection().await.map_err(E::from)?.query(msg, query).await
91 }
92}
93
94pub struct PinnedConnection {
101 conn: deadpool::managed::Object<ConnectionManager>,
102}
103
104impl PinnedConnection {
105 pub async fn transact<R, E, Q, M>(&self, msg: M, query: Q) -> std::result::Result<R, E>
108 where
109 Q: Send + for<'a> FnOnce(&'a mut SqliteConnection) -> std::result::Result<R, E> + 'static,
110 R: Send + 'static,
111 M: Send + ToString,
112 E: From<diesel::result::Error>,
113 E: From<DatabaseError>,
114 E: std::error::Error + Send + Sync + 'static,
115 {
116 let span = miden_node_tracing::Span::current();
117 self.conn
118 .interact(move |conn| {
119 let _guard = span.enter();
120 <_ as diesel::Connection>::transaction::<R, E, Q>(conn, query)
121 })
122 .await
123 .map_err(|err| E::from(DatabaseError::interact(&msg.to_string(), &err)))?
124 }
125
126 pub async fn query<R, E, Q, M>(&self, msg: M, query: Q) -> std::result::Result<R, E>
128 where
129 Q: Send + FnOnce(&mut SqliteConnection) -> std::result::Result<R, E> + 'static,
130 R: Send + 'static,
131 M: Send + ToString,
132 E: From<DatabaseError>,
133 E: std::error::Error + Send + Sync + 'static,
134 {
135 let span = miden_node_tracing::Span::current();
136 self.conn
137 .interact(move |conn| {
138 let _guard = span.enter();
139 query(conn)
140 })
141 .await
142 .map_err(|err| E::from(DatabaseError::interact(&msg.to_string(), &err)))?
143 }
144}