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