Skip to main content

a3s_orm/
executor.rs

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}