1use std::marker::PhantomData;
2
3use async_trait::async_trait;
4
5use crate::compiler::{CompiledQuery, Dialect};
6use crate::decode::{FromRow, Row};
7use crate::query::Query;
8use crate::Result;
9
10#[derive(Clone, Debug, Default, PartialEq, Eq)]
11pub struct ExecuteResult {
12 pub rows_affected: u64,
13}
14
15#[derive(Clone, Debug, PartialEq)]
16pub struct QueryResult<Row> {
17 pub rows: Vec<Row>,
18}
19
20#[async_trait]
21pub trait Executor: Send + Sync {
22 type Row: Send;
23 type Error: std::error::Error + Send + Sync + 'static;
24
25 async fn execute(
26 &self,
27 query: &CompiledQuery,
28 ) -> std::result::Result<ExecuteResult, Self::Error>;
29
30 async fn fetch_all(
31 &self,
32 query: &CompiledQuery,
33 ) -> std::result::Result<QueryResult<Self::Row>, Self::Error>;
34}
35
36#[async_trait]
37pub trait Transaction: Executor + Sized {
38 async fn commit(self) -> std::result::Result<(), Self::Error>;
39 async fn rollback(self) -> std::result::Result<(), Self::Error>;
40}
41
42#[async_trait]
43pub trait TransactionManager: Executor {
44 type Transaction: Transaction<Row = Self::Row, Error = Self::Error>;
45
46 async fn begin(&self) -> std::result::Result<Self::Transaction, Self::Error>;
47}
48
49#[derive(Debug)]
50pub struct Database<D, E> {
51 dialect: D,
52 executor: E,
53 marker: PhantomData<fn()>,
54}
55
56impl<D, E> Database<D, E>
57where
58 D: Dialect,
59 E: Executor,
60{
61 pub const fn new(dialect: D, executor: E) -> Self {
62 Self {
63 dialect,
64 executor,
65 marker: PhantomData,
66 }
67 }
68
69 pub fn dialect(&self) -> &D {
70 &self.dialect
71 }
72
73 pub fn executor(&self) -> &E {
74 &self.executor
75 }
76
77 pub fn compile<Q: Query>(&self, query: Q) -> Result<CompiledQuery> {
78 query.compile(&self.dialect)
79 }
80
81 pub async fn execute<Q: Query>(
82 &self,
83 query: Q,
84 ) -> std::result::Result<ExecuteResult, DatabaseError<E::Error>> {
85 let query = self.compile(query).map_err(DatabaseError::Build)?;
86 self.executor
87 .execute(&query)
88 .await
89 .map_err(DatabaseError::Execute)
90 }
91
92 pub async fn fetch_all<Q: Query>(
93 &self,
94 query: Q,
95 ) -> std::result::Result<QueryResult<E::Row>, DatabaseError<E::Error>> {
96 let query = self.compile(query).map_err(DatabaseError::Build)?;
97 self.executor
98 .fetch_all(&query)
99 .await
100 .map_err(DatabaseError::Execute)
101 }
102
103 pub async fn fetch_optional<Q: Query>(
104 &self,
105 query: Q,
106 ) -> std::result::Result<Option<E::Row>, DatabaseError<E::Error>> {
107 let result = self.fetch_all(query).await?;
108 exactly_optional(result.rows)
109 }
110
111 pub async fn fetch_one<Q: Query>(
112 &self,
113 query: Q,
114 ) -> std::result::Result<E::Row, DatabaseError<E::Error>> {
115 self.fetch_optional(query)
116 .await?
117 .ok_or(DatabaseError::NoRows)
118 }
119
120 pub async fn fetch_all_as<Q>(
121 &self,
122 query: Q,
123 ) -> std::result::Result<QueryResult<Q::Output>, DatabaseError<E::Error>>
124 where
125 Q: Query,
126 Q::Output: FromRow,
127 E::Row: Row,
128 {
129 let result = self.fetch_all(query).await?;
130 let rows = result
131 .rows
132 .iter()
133 .map(Q::Output::from_row)
134 .collect::<std::result::Result<Vec<_>, _>>()
135 .map_err(DatabaseError::Decode)?;
136 Ok(QueryResult { rows })
137 }
138
139 pub async fn fetch_optional_as<Q>(
140 &self,
141 query: Q,
142 ) -> std::result::Result<Option<Q::Output>, DatabaseError<E::Error>>
143 where
144 Q: Query,
145 Q::Output: FromRow,
146 E::Row: Row,
147 {
148 let result = self.fetch_all_as(query).await?;
149 exactly_optional(result.rows)
150 }
151
152 pub async fn fetch_one_as<Q>(
153 &self,
154 query: Q,
155 ) -> std::result::Result<Q::Output, DatabaseError<E::Error>>
156 where
157 Q: Query,
158 Q::Output: FromRow,
159 E::Row: Row,
160 {
161 self.fetch_optional_as(query)
162 .await?
163 .ok_or(DatabaseError::NoRows)
164 }
165
166 pub fn into_parts(self) -> (D, E) {
167 (self.dialect, self.executor)
168 }
169}
170
171#[derive(Debug, thiserror::Error)]
172pub enum DatabaseError<E>
173where
174 E: std::error::Error + 'static,
175{
176 #[error(transparent)]
177 Build(#[from] crate::Error),
178 #[error("database execution failed: {0}")]
179 Execute(E),
180 #[error("database row decoding failed: {0}")]
181 Decode(#[from] crate::DecodeError),
182 #[error("query returned no rows")]
183 NoRows,
184 #[error("query returned {actual} rows where at most one was expected")]
185 TooManyRows { actual: usize },
186}
187
188fn exactly_optional<Row, E>(rows: Vec<Row>) -> std::result::Result<Option<Row>, DatabaseError<E>>
189where
190 E: std::error::Error + Send + Sync + 'static,
191{
192 match rows.len() {
193 0 => Ok(None),
194 1 => Ok(rows.into_iter().next()),
195 actual => Err(DatabaseError::TooManyRows { actual }),
196 }
197}