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::{Expr, Record, SelectQuery};
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    PersistedRecord(String),
59}
60
61impl<E: std::error::Error + Send + Sync + 'static> std::fmt::Display for SqlExecutorError<E> {
62    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
63        match self {
64            SqlExecutorError::Compile(e) => write!(f, "SQL compile error: {}", e),
65            SqlExecutorError::Transport(e) => write!(f, "Transport error: {}", e),
66            SqlExecutorError::PersistedRecord(e) => write!(f, "Persisted record error: {}", e),
67        }
68    }
69}
70
71impl<E: std::error::Error + Send + Sync + 'static> std::error::Error for SqlExecutorError<E> {}
72
73#[derive(Clone)]
74pub struct SqlDataServiceExecutor<D, T, S> {
75    pub dialect: D,
76    pub transport: T,
77    pub schema_provider: S,
78}
79
80impl<D, T, S> SqlDataServiceExecutor<D, T, S> {
81    pub fn new(dialect: D, transport: T, schema_provider: S) -> Self {
82        Self {
83            dialect,
84            transport,
85            schema_provider,
86        }
87    }
88}
89
90impl<
91    D: SqlDialect + Send + Sync,
92    T: SqlTransport + Send + Sync,
93    S: teaql_data_service::SchemaProvider + Send + Sync,
94> DataServiceExecutor for SqlDataServiceExecutor<D, T, S>
95{
96    type Error = SqlExecutorError<T::Error>;
97
98    fn capabilities(&self) -> DataServiceCapabilities {
99        DataServiceCapabilities {
100            query: true,
101            mutation: true,
102            transaction: false, // Override if T implements SqlTransactionTransport
103            schema: false,
104            id_generation: false,
105            batch_mutation: true,
106            returning: false,
107        }
108    }
109}
110
111impl<
112    D: SqlDialect + Send + Sync,
113    T: SqlTransport + Send + Sync,
114    S: teaql_data_service::SchemaProvider + Send + Sync,
115> QueryExecutor for SqlDataServiceExecutor<D, T, S>
116{
117    fn query(
118        &self,
119        request: QueryRequest,
120    ) -> impl std::future::Future<Output = Result<QueryResult, Self::Error>> + Send {
121        async move {
122            let entity_desc = self
123                .schema_provider
124                .get_entity(&request.query.entity)
125                .ok_or_else(|| {
126                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(
127                        request.query.entity.clone(),
128                    ))
129                })?;
130
131            let compiled = self
132                .dialect
133                .compile_select(&entity_desc, &request.query)
134                .map_err(SqlExecutorError::Compile)?;
135            let start = SystemTime::now();
136            let rows = self
137                .transport
138                .fetch_all_sql(&compiled)
139                .await
140                .map_err(SqlExecutorError::Transport)?;
141            let end = SystemTime::now();
142
143            let metadata = ExecutionMetadata {
144                backend: "sql".to_string(),
145                operation: DataServiceOperation::Query,
146                started_at: start,
147                ended_at: end,
148                affected_rows: None,
149                result_count: Some(rows.len()),
150                trace_chain: request.trace_chain,
151                comment: request.comment,
152                backend_request_id: None,
153                parameterized_query: Some(compiled.sql.clone()),
154                params: compiled.params.clone(),
155                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
156            };
157
158            Ok(QueryResult { rows, metadata })
159        }
160    }
161}
162
163impl<
164    D: SqlDialect + Send + Sync,
165    T: SqlTransport + Send + Sync,
166    S: teaql_data_service::SchemaProvider + Send + Sync,
167> MutationExecutor for SqlDataServiceExecutor<D, T, S>
168{
169    fn mutate(
170        &self,
171        request: MutationRequest,
172    ) -> impl std::future::Future<Output = Result<MutationResult, Self::Error>> + Send {
173        async move {
174            let entity_name = match &request {
175                MutationRequest::Insert(cmd) => &cmd.entity,
176                MutationRequest::Update(cmd) => &cmd.entity,
177                MutationRequest::Delete(cmd) => &cmd.entity,
178                MutationRequest::Recover(cmd) => &cmd.entity,
179                MutationRequest::Batch(mutations) => {
180                    let mut total_affected = 0;
181                    let mut parameterized_queries = Vec::new();
182                    let mut params = Vec::new();
183                    let mut debug_queries = Vec::new();
184                    let start = SystemTime::now();
185                    for req in mutations {
186                        let res = Box::pin(self.mutate(req.clone())).await?;
187                        total_affected += res.affected_rows;
188                        if let Some(query) = res.metadata.parameterized_query {
189                            parameterized_queries.push(query);
190                        }
191                        params.extend(res.metadata.params);
192                        if let Some(query) = res.metadata.debug_query {
193                            debug_queries.push(query);
194                        }
195                    }
196                    let end = SystemTime::now();
197                    return Ok(MutationResult {
198                        affected_rows: total_affected,
199                        generated_values: Record::default(),
200                        persisted_record: None,
201                        metadata: ExecutionMetadata {
202                            backend: "sql".to_string(),
203                            operation: DataServiceOperation::Batch,
204                            started_at: start,
205                            ended_at: end,
206                            affected_rows: Some(total_affected),
207                            result_count: None,
208                            trace_chain: Vec::new(),
209                            comment: None,
210                            backend_request_id: None,
211                            parameterized_query: (!parameterized_queries.is_empty())
212                                .then(|| parameterized_queries.join("; ")),
213                            params,
214                            debug_query: (!debug_queries.is_empty())
215                                .then(|| debug_queries.join("; ")),
216                        },
217                    });
218                }
219            };
220
221            let entity_desc = self
222                .schema_provider
223                .get_entity(entity_name)
224                .ok_or_else(|| {
225                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(entity_name.clone()))
226                })?;
227
228            let compiled = match &request {
229                MutationRequest::Insert(cmd) => self
230                    .dialect
231                    .compile_insert(&entity_desc, cmd)
232                    .map_err(SqlExecutorError::Compile)?,
233                MutationRequest::Update(cmd) => self
234                    .dialect
235                    .compile_update(&entity_desc, cmd)
236                    .map_err(SqlExecutorError::Compile)?,
237                MutationRequest::Delete(cmd) => self
238                    .dialect
239                    .compile_delete(&entity_desc, cmd)
240                    .map_err(SqlExecutorError::Compile)?,
241                MutationRequest::Recover(cmd) => self
242                    .dialect
243                    .compile_recover(&entity_desc, cmd)
244                    .map_err(SqlExecutorError::Compile)?,
245                MutationRequest::Batch(_) => unreachable!(),
246            };
247
248            let start = SystemTime::now();
249            let affected_rows = self
250                .transport
251                .execute_sql(&compiled)
252                .await
253                .map_err(SqlExecutorError::Transport)?;
254            let end = SystemTime::now();
255
256            let operation = match &request {
257                MutationRequest::Insert(_) => DataServiceOperation::Insert,
258                MutationRequest::Update(_) => DataServiceOperation::Update,
259                MutationRequest::Delete(_) => DataServiceOperation::Delete,
260                MutationRequest::Recover(_) => DataServiceOperation::Recover,
261                MutationRequest::Batch(_) => DataServiceOperation::Batch,
262            };
263
264            let metadata = ExecutionMetadata {
265                backend: "sql".to_string(),
266                operation,
267                started_at: start,
268                ended_at: end,
269                affected_rows: Some(affected_rows),
270                result_count: None,
271                trace_chain: request.trace_chain().to_vec(),
272                comment: request.comment().map(|s| s.to_owned()),
273                backend_request_id: None,
274                parameterized_query: Some(compiled.sql.clone()),
275                params: compiled.params.clone(),
276                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
277            };
278
279            Ok(MutationResult {
280                affected_rows,
281                generated_values: Record::default(),
282                persisted_record: None,
283                metadata,
284            })
285        }
286    }
287}
288
289#[derive(Clone)]
290pub struct SqlDataServiceTransaction<'a, D, Tx: SqlTransport + SqlTransaction, S> {
291    pub dialect: &'a D,
292    pub transport: Tx,
293    pub schema_provider: &'a S,
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> DataServiceExecutor for SqlDataServiceTransaction<'a, D, Tx, S>
302{
303    type Error = SqlExecutorError<<Tx as SqlTransport>::Error>;
304
305    fn capabilities(&self) -> DataServiceCapabilities {
306        DataServiceCapabilities {
307            query: true,
308            mutation: true,
309            transaction: false,
310            schema: false,
311            id_generation: false,
312            batch_mutation: true,
313            returning: false,
314        }
315    }
316}
317
318impl<
319    'a,
320    D: SqlDialect + Send + Sync,
321    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
322    S: teaql_data_service::SchemaProvider + Send + Sync,
323> QueryExecutor for SqlDataServiceTransaction<'a, D, Tx, S>
324{
325    fn query(
326        &self,
327        request: QueryRequest,
328    ) -> impl std::future::Future<Output = Result<QueryResult, Self::Error>> + Send {
329        async move {
330            let entity_desc = self
331                .schema_provider
332                .get_entity(&request.query.entity)
333                .ok_or_else(|| {
334                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(
335                        request.query.entity.clone(),
336                    ))
337                })?;
338
339            let compiled = self
340                .dialect
341                .compile_select(&entity_desc, &request.query)
342                .map_err(SqlExecutorError::Compile)?;
343            let start = SystemTime::now();
344            let rows = self
345                .transport
346                .fetch_all_sql(&compiled)
347                .await
348                .map_err(SqlExecutorError::Transport)?;
349            let end = SystemTime::now();
350
351            let metadata = ExecutionMetadata {
352                backend: "sql".to_string(),
353                operation: DataServiceOperation::Query,
354                started_at: start,
355                ended_at: end,
356                affected_rows: None,
357                result_count: Some(rows.len()),
358                trace_chain: request.trace_chain,
359                comment: request.comment,
360                backend_request_id: None,
361                parameterized_query: Some(compiled.sql.clone()),
362                params: compiled.params.clone(),
363                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
364            };
365
366            Ok(QueryResult { rows, metadata })
367        }
368    }
369}
370
371impl<
372    'a,
373    D: SqlDialect + Send + Sync,
374    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
375    S: teaql_data_service::SchemaProvider + Send + Sync,
376> MutationExecutor for SqlDataServiceTransaction<'a, D, Tx, S>
377{
378    fn mutate(
379        &self,
380        request: MutationRequest,
381    ) -> impl std::future::Future<Output = Result<MutationResult, Self::Error>> + Send {
382        async move {
383            let entity_name = match &request {
384                MutationRequest::Insert(cmd) => &cmd.entity,
385                MutationRequest::Update(cmd) => &cmd.entity,
386                MutationRequest::Delete(cmd) => &cmd.entity,
387                MutationRequest::Recover(cmd) => &cmd.entity,
388                MutationRequest::Batch(mutations) => {
389                    let mut total_affected = 0;
390                    let mut parameterized_queries = Vec::new();
391                    let mut params = Vec::new();
392                    let mut debug_queries = Vec::new();
393                    let start = SystemTime::now();
394                    for req in mutations {
395                        let res = Box::pin(self.mutate(req.clone())).await?;
396                        total_affected += res.affected_rows;
397                        if let Some(query) = res.metadata.parameterized_query {
398                            parameterized_queries.push(query);
399                        }
400                        params.extend(res.metadata.params);
401                        if let Some(query) = res.metadata.debug_query {
402                            debug_queries.push(query);
403                        }
404                    }
405                    let end = SystemTime::now();
406                    return Ok(MutationResult {
407                        affected_rows: total_affected,
408                        generated_values: Record::default(),
409                        persisted_record: None,
410                        metadata: ExecutionMetadata {
411                            backend: "sql".to_string(),
412                            operation: DataServiceOperation::Batch,
413                            started_at: start,
414                            ended_at: end,
415                            affected_rows: Some(total_affected),
416                            result_count: None,
417                            trace_chain: Vec::new(),
418                            comment: None,
419                            backend_request_id: None,
420                            parameterized_query: (!parameterized_queries.is_empty())
421                                .then(|| parameterized_queries.join("; ")),
422                            params,
423                            debug_query: (!debug_queries.is_empty())
424                                .then(|| debug_queries.join("; ")),
425                        },
426                    });
427                }
428            };
429
430            let entity_desc = self
431                .schema_provider
432                .get_entity(entity_name)
433                .ok_or_else(|| {
434                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(entity_name.clone()))
435                })?;
436
437            let compiled = match &request {
438                MutationRequest::Insert(cmd) => self
439                    .dialect
440                    .compile_insert(&entity_desc, cmd)
441                    .map_err(SqlExecutorError::Compile)?,
442                MutationRequest::Update(cmd) => self
443                    .dialect
444                    .compile_update(&entity_desc, cmd)
445                    .map_err(SqlExecutorError::Compile)?,
446                MutationRequest::Delete(cmd) => self
447                    .dialect
448                    .compile_delete(&entity_desc, cmd)
449                    .map_err(SqlExecutorError::Compile)?,
450                MutationRequest::Recover(cmd) => self
451                    .dialect
452                    .compile_recover(&entity_desc, cmd)
453                    .map_err(SqlExecutorError::Compile)?,
454                MutationRequest::Batch(_) => unreachable!("batch handled above"),
455            };
456
457            let start = SystemTime::now();
458            let affected_rows = self
459                .transport
460                .execute_sql(&compiled)
461                .await
462                .map_err(SqlExecutorError::Transport)?;
463            let end = SystemTime::now();
464
465            let operation = match &request {
466                MutationRequest::Insert(_) => DataServiceOperation::Insert,
467                MutationRequest::Update(_) => DataServiceOperation::Update,
468                MutationRequest::Delete(_) => DataServiceOperation::Delete,
469                MutationRequest::Recover(_) => DataServiceOperation::Recover,
470                MutationRequest::Batch(_) => DataServiceOperation::Batch,
471            };
472
473            let persisted_id = match &request {
474                MutationRequest::Insert(cmd) => cmd.values.get("id").cloned(),
475                MutationRequest::Update(cmd) => Some(cmd.id.clone()),
476                MutationRequest::Delete(cmd) if cmd.soft_delete => Some(cmd.id.clone()),
477                MutationRequest::Recover(cmd) => Some(cmd.id.clone()),
478                MutationRequest::Delete(_) | MutationRequest::Batch(_) => None,
479            };
480            let persisted_record = if affected_rows == 1 {
481                if let Some(id) = persisted_id {
482                    let query = SelectQuery::new(entity_name.clone()).filter(Expr::eq("id", id));
483                    let compiled_readback = self
484                        .dialect
485                        .compile_select(&entity_desc, &query)
486                        .map_err(SqlExecutorError::Compile)?;
487                    let mut rows = self
488                        .transport
489                        .fetch_all_sql(&compiled_readback)
490                        .await
491                        .map_err(SqlExecutorError::Transport)?;
492                    if rows.len() != 1 {
493                        return Err(SqlExecutorError::PersistedRecord(format!(
494                            "persisted {entity_name} record could not be read back"
495                        )));
496                    }
497                    rows.pop()
498                } else {
499                    None
500                }
501            } else {
502                None
503            };
504
505            let metadata = ExecutionMetadata {
506                backend: "sql".to_string(),
507                operation,
508                started_at: start,
509                ended_at: end,
510                affected_rows: Some(affected_rows),
511                result_count: None,
512                trace_chain: request.trace_chain().to_vec(),
513                comment: request.comment().map(|s| s.to_owned()),
514                backend_request_id: None,
515                parameterized_query: Some(compiled.sql.clone()),
516                params: compiled.params.clone(),
517                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
518            };
519
520            Ok(MutationResult {
521                affected_rows,
522                generated_values: Record::default(),
523                persisted_record,
524                metadata,
525            })
526        }
527    }
528}
529
530impl<
531    'a,
532    D: SqlDialect + Send + Sync,
533    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
534    S: teaql_data_service::SchemaProvider + Send + Sync,
535> teaql_data_service::Transaction for SqlDataServiceTransaction<'a, D, Tx, S>
536{
537    type Error = SqlExecutorError<<Tx as SqlTransport>::Error>;
538
539    fn commit(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send {
540        async move {
541            self.transport
542                .commit_sql()
543                .await
544                .map_err(SqlExecutorError::Transport)
545        }
546    }
547
548    fn rollback(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send {
549        async move {
550            self.transport
551                .rollback_sql()
552                .await
553                .map_err(SqlExecutorError::Transport)
554        }
555    }
556}
557
558impl<
559    D: SqlDialect + Send + Sync,
560    T: SqlTransactionTransport + Send + Sync,
561    S: teaql_data_service::SchemaProvider + Send + Sync,
562> teaql_data_service::TransactionExecutor for SqlDataServiceExecutor<D, T, S>
563{
564    type Tx<'a>
565        = SqlDataServiceTransaction<'a, D, T::Tx<'a>, S>
566    where
567        Self: 'a;
568
569    fn begin(&self) -> impl std::future::Future<Output = Result<Self::Tx<'_>, Self::Error>> + Send {
570        async move {
571            let tx = self
572                .transport
573                .begin_sql()
574                .await
575                .map_err(SqlExecutorError::Transport)?;
576            Ok(SqlDataServiceTransaction {
577                dialect: &self.dialect,
578                transport: tx,
579                schema_provider: &self.schema_provider,
580            })
581        }
582    }
583}
584
585impl<
586    D: SqlDialect + Send + Sync,
587    T: StreamingSqlTransport + Send + Sync,
588    S: teaql_data_service::SchemaProvider + Send + Sync,
589> teaql_data_service::StreamQueryExecutor for SqlDataServiceExecutor<D, T, S>
590{
591    fn query_stream(
592        &self,
593        request: teaql_data_service::QueryRequest,
594        chunk_size: usize,
595    ) -> teaql_data_service::QueryStream<'_, Self::Error> {
596        use futures_util::StreamExt;
597        let entity = match self.schema_provider.get_entity(&request.query.entity) {
598            Some(entity) => entity,
599            None => {
600                return Box::pin(futures_util::stream::once(async {
601                    Err(SqlExecutorError::Compile(SqlCompileError::UnknownEntity(
602                        request.query.entity,
603                    )))
604                }));
605            }
606        };
607        match self.dialect.compile_select(&entity, &request.query) {
608            Ok(compiled) => Box::pin(
609                self.transport
610                    .stream_sql(compiled, chunk_size)
611                    .map(|r| r.map_err(SqlExecutorError::Transport)),
612            ),
613            Err(error) => Box::pin(futures_util::stream::once(async {
614                Err(SqlExecutorError::Compile(error))
615            })),
616        }
617    }
618}