Skip to main content

teaql_sql/
executor.rs

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