Skip to main content

teaql_sql/
executor.rs

1#![allow(clippy::manual_async_fn)]
2#![allow(async_fn_in_trait)]
3
4use std::time::SystemTime;
5use teaql_core::Record;
6use teaql_data_service::{
7    DataServiceCapabilities, DataServiceExecutor, DataServiceOperation, ExecutionMetadata,
8    MutationExecutor, MutationRequest, MutationResult, QueryExecutor, QueryRequest, QueryResult,
9};
10
11use crate::{CompiledQuery, SqlCompileError, SqlDialect};
12
13pub trait SqlTransport: Send + Sync {
14    type Error: std::error::Error + Send + Sync + 'static;
15
16    fn fetch_all_sql(
17        &self,
18        query: &CompiledQuery,
19    ) -> impl std::future::Future<Output = Result<Vec<Record>, Self::Error>> + Send;
20    fn execute_sql(
21        &self,
22        query: &CompiledQuery,
23    ) -> impl std::future::Future<Output = Result<u64, Self::Error>> + Send;
24}
25
26pub trait StreamingSqlTransport: SqlTransport {
27    fn stream_sql(
28        &self,
29        query: CompiledQuery,
30        chunk_size: usize,
31    ) -> teaql_data_service::QueryStream<'_, Self::Error>;
32}
33
34pub trait SqlTransactionTransport: SqlTransport {
35    type Tx<'a>: SqlTransport<Error = Self::Error>
36        + SqlTransaction<Error = Self::Error>
37        + Send
38        + Sync
39        + 'a
40    where
41        Self: 'a;
42
43    fn begin_sql(
44        &self,
45    ) -> impl std::future::Future<Output = Result<Self::Tx<'_>, Self::Error>> + Send;
46}
47
48pub trait SqlTransaction {
49    type Error: std::error::Error + Send + Sync + 'static;
50    fn commit_sql(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send;
51    fn rollback_sql(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send;
52}
53
54#[derive(Debug)]
55pub enum SqlExecutorError<E: std::error::Error + Send + Sync + 'static> {
56    Compile(SqlCompileError),
57    Transport(E),
58}
59
60impl<E: std::error::Error + Send + Sync + 'static> std::fmt::Display for SqlExecutorError<E> {
61    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
62        match self {
63            SqlExecutorError::Compile(e) => write!(f, "SQL compile error: {}", e),
64            SqlExecutorError::Transport(e) => write!(f, "Transport error: {}", e),
65        }
66    }
67}
68
69impl<E: std::error::Error + Send + Sync + 'static> std::error::Error for SqlExecutorError<E> {}
70
71#[derive(Clone)]
72pub struct SqlDataServiceExecutor<D, T, S> {
73    pub dialect: D,
74    pub transport: T,
75    pub schema_provider: S,
76}
77
78impl<D, T, S> SqlDataServiceExecutor<D, T, S> {
79    pub fn new(dialect: D, transport: T, schema_provider: S) -> Self {
80        Self {
81            dialect,
82            transport,
83            schema_provider,
84        }
85    }
86}
87
88impl<
89    D: SqlDialect + Send + Sync,
90    T: SqlTransport + Send + Sync,
91    S: teaql_data_service::SchemaProvider + Send + Sync,
92> DataServiceExecutor for SqlDataServiceExecutor<D, T, S>
93{
94    type Error = SqlExecutorError<T::Error>;
95
96    fn capabilities(&self) -> DataServiceCapabilities {
97        DataServiceCapabilities {
98            query: true,
99            mutation: true,
100            transaction: false, // Override if T implements SqlTransactionTransport
101            schema: false,
102            id_generation: false,
103            batch_mutation: true,
104            returning: false,
105        }
106    }
107}
108
109impl<
110    D: SqlDialect + Send + Sync,
111    T: SqlTransport + Send + Sync,
112    S: teaql_data_service::SchemaProvider + Send + Sync,
113> QueryExecutor for SqlDataServiceExecutor<D, T, S>
114{
115    fn query(
116        &self,
117        request: QueryRequest,
118    ) -> impl std::future::Future<Output = Result<QueryResult, Self::Error>> + Send {
119        async move {
120            let entity_desc = self
121                .schema_provider
122                .get_entity(&request.query.entity)
123                .ok_or_else(|| {
124                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(
125                        request.query.entity.clone(),
126                    ))
127                })?;
128
129            let compiled = self
130                .dialect
131                .compile_select(&entity_desc, &request.query)
132                .map_err(SqlExecutorError::Compile)?;
133            let start = SystemTime::now();
134            let rows = self
135                .transport
136                .fetch_all_sql(&compiled)
137                .await
138                .map_err(SqlExecutorError::Transport)?;
139            let end = SystemTime::now();
140
141            let metadata = ExecutionMetadata {
142                backend: "sql".to_string(),
143                operation: DataServiceOperation::Query,
144                started_at: start,
145                ended_at: end,
146                affected_rows: None,
147                result_count: Some(rows.len()),
148                trace_chain: request.trace_chain,
149                comment: request.comment,
150                backend_request_id: None,
151                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
152            };
153
154            Ok(QueryResult { rows, metadata })
155        }
156    }
157}
158
159impl<
160    D: SqlDialect + Send + Sync,
161    T: SqlTransport + Send + Sync,
162    S: teaql_data_service::SchemaProvider + Send + Sync,
163> MutationExecutor for SqlDataServiceExecutor<D, T, S>
164{
165    fn mutate(
166        &self,
167        request: MutationRequest,
168    ) -> impl std::future::Future<Output = Result<MutationResult, Self::Error>> + Send {
169        async move {
170            let entity_name = match &request {
171                MutationRequest::Insert(cmd) => &cmd.entity,
172                MutationRequest::Update(cmd) => &cmd.entity,
173                MutationRequest::Delete(cmd) => &cmd.entity,
174                MutationRequest::Recover(cmd) => &cmd.entity,
175                MutationRequest::Batch(mutations) => {
176                    let mut total_affected = 0;
177                    let start = SystemTime::now();
178                    for req in mutations {
179                        let res = Box::pin(self.mutate(req.clone())).await?;
180                        total_affected += res.affected_rows;
181                    }
182                    let end = SystemTime::now();
183                    return Ok(MutationResult {
184                        affected_rows: total_affected,
185                        generated_values: Record::default(),
186                        metadata: ExecutionMetadata {
187                            backend: "sql".to_string(),
188                            operation: DataServiceOperation::Batch,
189                            started_at: start,
190                            ended_at: end,
191                            affected_rows: Some(total_affected),
192                            result_count: None,
193                            trace_chain: Vec::new(),
194                            comment: None,
195                            backend_request_id: None,
196                            debug_query: None,
197                        },
198                    });
199                }
200            };
201
202            let entity_desc = self
203                .schema_provider
204                .get_entity(entity_name)
205                .ok_or_else(|| {
206                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(entity_name.clone()))
207                })?;
208
209            let compiled = match &request {
210                MutationRequest::Insert(cmd) => self
211                    .dialect
212                    .compile_insert(&entity_desc, cmd)
213                    .map_err(SqlExecutorError::Compile)?,
214                MutationRequest::Update(cmd) => self
215                    .dialect
216                    .compile_update(&entity_desc, cmd)
217                    .map_err(SqlExecutorError::Compile)?,
218                MutationRequest::Delete(cmd) => self
219                    .dialect
220                    .compile_delete(&entity_desc, cmd)
221                    .map_err(SqlExecutorError::Compile)?,
222                MutationRequest::Recover(cmd) => self
223                    .dialect
224                    .compile_recover(&entity_desc, cmd)
225                    .map_err(SqlExecutorError::Compile)?,
226                MutationRequest::Batch(_) => unreachable!(),
227            };
228
229            let start = SystemTime::now();
230            let affected_rows = self
231                .transport
232                .execute_sql(&compiled)
233                .await
234                .map_err(SqlExecutorError::Transport)?;
235            let end = SystemTime::now();
236
237            let operation = match &request {
238                MutationRequest::Insert(_) => DataServiceOperation::Insert,
239                MutationRequest::Update(_) => DataServiceOperation::Update,
240                MutationRequest::Delete(_) => DataServiceOperation::Delete,
241                MutationRequest::Recover(_) => DataServiceOperation::Recover,
242                MutationRequest::Batch(_) => DataServiceOperation::Batch,
243            };
244
245            let metadata = ExecutionMetadata {
246                backend: "sql".to_string(),
247                operation,
248                started_at: start,
249                ended_at: end,
250                affected_rows: Some(affected_rows),
251                result_count: None,
252                trace_chain: request.trace_chain().to_vec(),
253                comment: request.comment().map(|s| s.to_owned()),
254                backend_request_id: None,
255                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
256            };
257
258            Ok(MutationResult {
259                affected_rows,
260                generated_values: Record::default(),
261                metadata,
262            })
263        }
264    }
265}
266
267#[derive(Clone)]
268pub struct SqlDataServiceTransaction<'a, D, Tx: SqlTransport + SqlTransaction, S> {
269    pub dialect: &'a D,
270    pub transport: Tx,
271    pub schema_provider: &'a S,
272}
273
274impl<
275    'a,
276    D: SqlDialect + Send + Sync,
277    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
278    S: teaql_data_service::SchemaProvider + Send + Sync,
279> DataServiceExecutor for SqlDataServiceTransaction<'a, D, Tx, S>
280{
281    type Error = SqlExecutorError<<Tx as SqlTransport>::Error>;
282
283    fn capabilities(&self) -> DataServiceCapabilities {
284        DataServiceCapabilities {
285            query: true,
286            mutation: true,
287            transaction: false,
288            schema: false,
289            id_generation: false,
290            batch_mutation: true,
291            returning: false,
292        }
293    }
294}
295
296impl<
297    'a,
298    D: SqlDialect + Send + Sync,
299    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
300    S: teaql_data_service::SchemaProvider + Send + Sync,
301> QueryExecutor for SqlDataServiceTransaction<'a, D, Tx, S>
302{
303    fn query(
304        &self,
305        request: QueryRequest,
306    ) -> impl std::future::Future<Output = Result<QueryResult, Self::Error>> + Send {
307        async move {
308            let entity_desc = self
309                .schema_provider
310                .get_entity(&request.query.entity)
311                .ok_or_else(|| {
312                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(
313                        request.query.entity.clone(),
314                    ))
315                })?;
316
317            let compiled = self
318                .dialect
319                .compile_select(&entity_desc, &request.query)
320                .map_err(SqlExecutorError::Compile)?;
321            let start = SystemTime::now();
322            let rows = self
323                .transport
324                .fetch_all_sql(&compiled)
325                .await
326                .map_err(SqlExecutorError::Transport)?;
327            let end = SystemTime::now();
328
329            let metadata = ExecutionMetadata {
330                backend: "sql".to_string(),
331                operation: DataServiceOperation::Query,
332                started_at: start,
333                ended_at: end,
334                affected_rows: None,
335                result_count: Some(rows.len()),
336                trace_chain: request.trace_chain,
337                comment: request.comment,
338                backend_request_id: None,
339                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
340            };
341
342            Ok(QueryResult { rows, metadata })
343        }
344    }
345}
346
347impl<
348    'a,
349    D: SqlDialect + Send + Sync,
350    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
351    S: teaql_data_service::SchemaProvider + Send + Sync,
352> MutationExecutor for SqlDataServiceTransaction<'a, D, Tx, S>
353{
354    fn mutate(
355        &self,
356        request: MutationRequest,
357    ) -> impl std::future::Future<Output = Result<MutationResult, Self::Error>> + Send {
358        async move {
359            let entity_name = match &request {
360                MutationRequest::Insert(cmd) => &cmd.entity,
361                MutationRequest::Update(cmd) => &cmd.entity,
362                MutationRequest::Delete(cmd) => &cmd.entity,
363                MutationRequest::Recover(cmd) => &cmd.entity,
364                MutationRequest::Batch(mutations) => {
365                    let mut total_affected = 0;
366                    let start = SystemTime::now();
367                    for req in mutations {
368                        let res = Box::pin(self.mutate(req.clone())).await?;
369                        total_affected += res.affected_rows;
370                    }
371                    let end = SystemTime::now();
372                    return Ok(MutationResult {
373                        affected_rows: total_affected,
374                        generated_values: Record::default(),
375                        metadata: ExecutionMetadata {
376                            backend: "sql".to_string(),
377                            operation: DataServiceOperation::Batch,
378                            started_at: start,
379                            ended_at: end,
380                            affected_rows: Some(total_affected),
381                            result_count: None,
382                            trace_chain: Vec::new(),
383                            comment: None,
384                            backend_request_id: None,
385                            debug_query: None,
386                        },
387                    });
388                }
389            };
390
391            let entity_desc = self
392                .schema_provider
393                .get_entity(entity_name)
394                .ok_or_else(|| {
395                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(entity_name.clone()))
396                })?;
397
398            let compiled = match &request {
399                MutationRequest::Insert(cmd) => self
400                    .dialect
401                    .compile_insert(&entity_desc, cmd)
402                    .map_err(SqlExecutorError::Compile)?,
403                MutationRequest::Update(cmd) => self
404                    .dialect
405                    .compile_update(&entity_desc, cmd)
406                    .map_err(SqlExecutorError::Compile)?,
407                MutationRequest::Delete(cmd) => self
408                    .dialect
409                    .compile_delete(&entity_desc, cmd)
410                    .map_err(SqlExecutorError::Compile)?,
411                MutationRequest::Recover(cmd) => self
412                    .dialect
413                    .compile_recover(&entity_desc, cmd)
414                    .map_err(SqlExecutorError::Compile)?,
415                MutationRequest::Batch(_) => unreachable!("batch handled above"),
416            };
417
418            let start = SystemTime::now();
419            let affected_rows = self
420                .transport
421                .execute_sql(&compiled)
422                .await
423                .map_err(SqlExecutorError::Transport)?;
424            let end = SystemTime::now();
425
426            let operation = match &request {
427                MutationRequest::Insert(_) => DataServiceOperation::Insert,
428                MutationRequest::Update(_) => DataServiceOperation::Update,
429                MutationRequest::Delete(_) => DataServiceOperation::Delete,
430                MutationRequest::Recover(_) => DataServiceOperation::Recover,
431                MutationRequest::Batch(_) => DataServiceOperation::Batch,
432            };
433
434            let metadata = ExecutionMetadata {
435                backend: "sql".to_string(),
436                operation,
437                started_at: start,
438                ended_at: end,
439                affected_rows: Some(affected_rows),
440                result_count: None,
441                trace_chain: request.trace_chain().to_vec(),
442                comment: request.comment().map(|s| s.to_owned()),
443                backend_request_id: None,
444                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
445            };
446
447            Ok(MutationResult {
448                affected_rows,
449                generated_values: Record::default(),
450                metadata,
451            })
452        }
453    }
454}
455
456impl<
457    'a,
458    D: SqlDialect + Send + Sync,
459    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
460    S: teaql_data_service::SchemaProvider + Send + Sync,
461> teaql_data_service::Transaction for SqlDataServiceTransaction<'a, D, Tx, S>
462{
463    type Error = SqlExecutorError<<Tx as SqlTransport>::Error>;
464
465    fn commit(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send {
466        async move {
467            self.transport
468                .commit_sql()
469                .await
470                .map_err(SqlExecutorError::Transport)
471        }
472    }
473
474    fn rollback(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send {
475        async move {
476            self.transport
477                .rollback_sql()
478                .await
479                .map_err(SqlExecutorError::Transport)
480        }
481    }
482}
483
484impl<
485    D: SqlDialect + Send + Sync,
486    T: SqlTransactionTransport + Send + Sync,
487    S: teaql_data_service::SchemaProvider + Send + Sync,
488> teaql_data_service::TransactionExecutor for SqlDataServiceExecutor<D, T, S>
489{
490    type Tx<'a>
491        = SqlDataServiceTransaction<'a, D, T::Tx<'a>, S>
492    where
493        Self: 'a;
494
495    fn begin(&self) -> impl std::future::Future<Output = Result<Self::Tx<'_>, Self::Error>> + Send {
496        async move {
497            let tx = self
498                .transport
499                .begin_sql()
500                .await
501                .map_err(SqlExecutorError::Transport)?;
502            Ok(SqlDataServiceTransaction {
503                dialect: &self.dialect,
504                transport: tx,
505                schema_provider: &self.schema_provider,
506            })
507        }
508    }
509}
510
511impl<
512    D: SqlDialect + Send + Sync,
513    T: StreamingSqlTransport + Send + Sync,
514    S: teaql_data_service::SchemaProvider + Send + Sync,
515> teaql_data_service::StreamQueryExecutor for SqlDataServiceExecutor<D, T, S>
516{
517    fn query_stream(
518        &self,
519        request: teaql_data_service::QueryRequest,
520        chunk_size: usize,
521    ) -> teaql_data_service::QueryStream<'_, Self::Error> {
522        use futures_util::StreamExt;
523        let entity = match self.schema_provider.get_entity(&request.query.entity) {
524            Some(entity) => entity,
525            None => {
526                return Box::pin(futures_util::stream::once(async {
527                    Err(SqlExecutorError::Compile(SqlCompileError::UnknownEntity(
528                        request.query.entity,
529                    )))
530                }));
531            }
532        };
533        match self.dialect.compile_select(&entity, &request.query) {
534            Ok(compiled) => Box::pin(
535                self.transport
536                    .stream_sql(compiled, chunk_size)
537                    .map(|r| r.map_err(SqlExecutorError::Transport)),
538            ),
539            Err(error) => Box::pin(futures_util::stream::once(async {
540                Err(SqlExecutorError::Compile(error))
541            })),
542        }
543    }
544}