Skip to main content

teaql_sql/
executor.rs

1#![allow(clippy::manual_async_fn)]
2#![allow(async_fn_in_trait)]
3
4use std::collections::HashMap;
5use std::sync::{Arc, RwLock};
6use std::time::SystemTime;
7use teaql_core::{EntityDescriptor, Expr, Record, SelectQuery, Value};
8use teaql_data_service::{
9    DataServiceCapabilities, DataServiceExecutor, DataServiceOperation, ExecutionMetadata,
10    MutationExecutor, MutationRequest, MutationResult, QueryExecutor, QueryRequest, QueryResult,
11};
12
13use crate::{CompiledQuery, SqlCompileError, SqlDialect};
14
15pub trait SqlTransport: Send + Sync {
16    type Error: std::error::Error + Send + Sync + 'static;
17
18    fn fetch_all_sql(
19        &self,
20        query: &CompiledQuery,
21    ) -> impl std::future::Future<Output = Result<Vec<Record>, Self::Error>> + Send;
22    fn execute_sql(
23        &self,
24        query: &CompiledQuery,
25    ) -> impl std::future::Future<Output = Result<u64, Self::Error>> + Send;
26}
27
28pub trait StreamingSqlTransport: SqlTransport {
29    fn stream_sql(
30        &self,
31        query: CompiledQuery,
32        chunk_size: usize,
33    ) -> teaql_data_service::QueryStream<'_, Self::Error>;
34}
35
36pub trait SqlTransactionTransport: SqlTransport {
37    type Tx<'a>: SqlTransport<Error = Self::Error>
38        + SqlTransaction<Error = Self::Error>
39        + Send
40        + Sync
41        + 'a
42    where
43        Self: 'a;
44
45    fn begin_sql(
46        &self,
47    ) -> impl std::future::Future<Output = Result<Self::Tx<'_>, Self::Error>> + Send;
48}
49
50pub trait SqlTransaction {
51    type Error: std::error::Error + Send + Sync + 'static;
52    fn commit_sql(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send;
53    fn rollback_sql(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send;
54}
55
56#[derive(Debug)]
57pub enum SqlExecutorError<E: std::error::Error + Send + Sync + 'static> {
58    Compile(SqlCompileError),
59    Transport(E),
60    PersistedRecord(String),
61}
62
63impl<E: std::error::Error + Send + Sync + 'static> std::fmt::Display for SqlExecutorError<E> {
64    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
65        match self {
66            SqlExecutorError::Compile(e) => write!(f, "SQL compile error: {}", e),
67            SqlExecutorError::Transport(e) => write!(f, "Transport error: {}", e),
68            SqlExecutorError::PersistedRecord(e) => write!(f, "Persisted record error: {}", e),
69        }
70    }
71}
72
73impl<E: std::error::Error + Send + Sync + 'static> std::error::Error for SqlExecutorError<E> {}
74
75#[derive(Clone)]
76pub struct SqlDataServiceExecutor<D, T, S> {
77    pub dialect: D,
78    pub transport: T,
79    pub schema_provider: S,
80    descriptor_cache: Arc<RwLock<HashMap<String, Arc<teaql_core::EntityDescriptor>>>>,
81    select_plan_cache: Arc<RwLock<Vec<(SelectQuery, String)>>>,
82}
83
84impl<D, T, S> SqlDataServiceExecutor<D, T, S> {
85    pub fn new(dialect: D, transport: T, schema_provider: S) -> Self {
86        Self {
87            dialect,
88            transport,
89            schema_provider,
90            descriptor_cache: Arc::new(RwLock::new(HashMap::new())),
91            select_plan_cache: Arc::new(RwLock::new(Vec::new())),
92        }
93    }
94}
95
96impl<D, T, S> SqlDataServiceExecutor<D, T, S>
97where
98    D: SqlDialect,
99    S: teaql_data_service::SchemaProvider,
100{
101    fn compile_select_cached(
102        &self,
103        entity: &EntityDescriptor,
104        query: &SelectQuery,
105    ) -> Result<CompiledQuery, SqlCompileError> {
106        compile_select_with_cache(&self.dialect, &self.select_plan_cache, entity, query)
107    }
108}
109
110fn compile_select_with_cache<D: SqlDialect>(
111    dialect: &D,
112    plan_cache: &RwLock<Vec<(SelectQuery, String)>>,
113    entity: &EntityDescriptor,
114    query: &SelectQuery,
115) -> Result<CompiledQuery, SqlCompileError> {
116    if let Ok(cache) = plan_cache.read()
117        && let Some((_, sql)) = cache
118            .iter()
119            .find(|(candidate, _)| select_plan_matches(candidate, query))
120    {
121        return Ok(CompiledQuery {
122            sql: sql.clone(),
123            params: collect_select_params(entity, query),
124            comment: query.comment.clone(),
125        });
126    }
127
128    let key = select_plan_key(query);
129    let compiled = dialect.compile_select(entity, query)?;
130    if let Ok(mut cache) = plan_cache.write() {
131        if cache.len() >= 256 {
132            cache.remove(0);
133        }
134        if !cache
135            .iter()
136            .any(|(candidate, _)| select_plan_matches(candidate, query))
137        {
138            cache.push((key, compiled.sql.clone()));
139        }
140    }
141    Ok(compiled)
142}
143
144fn select_plan_matches(key: &SelectQuery, query: &SelectQuery) -> bool {
145    key.hard_limit == query.hard_limit
146        && key.entity == query.entity
147        && key.projection == query.projection
148        && key.expr_projection.len() == query.expr_projection.len()
149        && key
150            .expr_projection
151            .iter()
152            .zip(&query.expr_projection)
153            .all(|(left, right)| {
154                left.alias == right.alias && expr_plan_matches(&left.expr, &right.expr)
155            })
156        && key.search_with_text.is_some() == query.search_with_text.is_some()
157        && optional_expr_plan_matches(key.filter.as_ref(), query.filter.as_ref())
158        && optional_expr_plan_matches(key.having.as_ref(), query.having.as_ref())
159        && key.order_by.len() == query.order_by.len()
160        && key
161            .order_by
162            .iter()
163            .zip(&query.order_by)
164            .all(|(left, right)| {
165                left.field == right.field
166                    && left.direction == right.direction
167                    && optional_expr_plan_matches(left.expr.as_ref(), right.expr.as_ref())
168            })
169        && key.slice == query.slice
170        && key.partition_by == query.partition_by
171        && key.aggregates == query.aggregates
172        && key.group_by == query.group_by
173        && key.relations == query.relations
174        && key.aggregation_cache == query.aggregation_cache
175        && key.raw_sql == query.raw_sql
176        && key.raw_sql_search_criteria == query.raw_sql_search_criteria
177        && key.dynamic_properties == query.dynamic_properties
178        && key.raw_projections == query.raw_projections
179        && key.object_group_bys == query.object_group_bys
180        && key.child_enhancements == query.child_enhancements
181        && key.stream_config == query.stream_config
182        && key.continuous_page_fetch == query.continuous_page_fetch
183}
184
185fn optional_expr_plan_matches(key: Option<&Expr>, query: Option<&Expr>) -> bool {
186    match (key, query) {
187        (Some(key), Some(query)) => expr_plan_matches(key, query),
188        (None, None) => true,
189        _ => false,
190    }
191}
192
193fn expr_plan_matches(key: &Expr, query: &Expr) -> bool {
194    match (key, query) {
195        (Expr::Column(left), Expr::Column(right)) => left == right,
196        (Expr::Value(Value::List(left)), Expr::Value(Value::List(right))) => {
197            left.len() == right.len()
198        }
199        (Expr::Value(Value::List(_)), Expr::Value(_))
200        | (Expr::Value(_), Expr::Value(Value::List(_))) => false,
201        (Expr::Value(_), Expr::Value(_)) => true,
202        (
203            Expr::Function {
204                function: left_function,
205                args: left_args,
206            },
207            Expr::Function {
208                function: right_function,
209                args: right_args,
210            },
211        ) => left_function == right_function && expr_slice_plan_matches(left_args, right_args),
212        (
213            Expr::Binary {
214                left: left_left,
215                op: left_op,
216                right: left_right,
217            },
218            Expr::Binary {
219                left: right_left,
220                op: right_op,
221                right: right_right,
222            },
223        ) => {
224            left_op == right_op
225                && expr_plan_matches(left_left, right_left)
226                && expr_plan_matches(left_right, right_right)
227        }
228        (
229            Expr::SubQuery {
230                left: left_expr,
231                op: left_op,
232                entity: left_entity,
233                query: left_query,
234            },
235            Expr::SubQuery {
236                left: right_expr,
237                op: right_op,
238                entity: right_entity,
239                query: right_query,
240            },
241        ) => {
242            left_op == right_op
243                && left_entity == right_entity
244                && expr_plan_matches(left_expr, right_expr)
245                && select_plan_matches(left_query, right_query)
246        }
247        (
248            Expr::Between {
249                expr: left_expr,
250                lower: left_lower,
251                upper: left_upper,
252            },
253            Expr::Between {
254                expr: right_expr,
255                lower: right_lower,
256                upper: right_upper,
257            },
258        ) => {
259            expr_plan_matches(left_expr, right_expr)
260                && expr_plan_matches(left_lower, right_lower)
261                && expr_plan_matches(left_upper, right_upper)
262        }
263        (Expr::IsNull(left), Expr::IsNull(right))
264        | (Expr::IsNotNull(left), Expr::IsNotNull(right))
265        | (Expr::Not(left), Expr::Not(right)) => expr_plan_matches(left, right),
266        (Expr::And(left), Expr::And(right)) | (Expr::Or(left), Expr::Or(right)) => {
267            expr_slice_plan_matches(left, right)
268        }
269        _ => false,
270    }
271}
272
273fn expr_slice_plan_matches(left: &[Expr], right: &[Expr]) -> bool {
274    left.len() == right.len()
275        && left
276            .iter()
277            .zip(right)
278            .all(|(left, right)| expr_plan_matches(left, right))
279}
280
281fn select_plan_key(query: &SelectQuery) -> SelectQuery {
282    let mut key = query.clone();
283    key.comment = None;
284    key.trace_chain.clear();
285    if key.search_with_text.is_some() {
286        key.search_with_text = Some(String::new());
287    }
288    for projection in &mut key.expr_projection {
289        normalize_expr_values(&mut projection.expr);
290    }
291    if let Some(expr) = &mut key.filter {
292        normalize_expr_values(expr);
293    }
294    if let Some(expr) = &mut key.having {
295        normalize_expr_values(expr);
296    }
297    for order in &mut key.order_by {
298        if let Some(expr) = &mut order.expr {
299            normalize_expr_values(expr);
300        }
301    }
302    key
303}
304
305fn normalize_expr_values(expr: &mut Expr) {
306    match expr {
307        Expr::Value(Value::List(values)) => {
308            for value in values {
309                *value = Value::Null;
310            }
311        }
312        Expr::Value(value) => *value = Value::Null,
313        Expr::Function { args, .. } | Expr::And(args) | Expr::Or(args) => {
314            for arg in args {
315                normalize_expr_values(arg);
316            }
317        }
318        Expr::Binary { left, right, .. } => {
319            normalize_expr_values(left);
320            normalize_expr_values(right);
321        }
322        Expr::SubQuery { left, query, .. } => {
323            normalize_expr_values(left);
324            **query = select_plan_key(query);
325        }
326        Expr::Between { expr, lower, upper } => {
327            normalize_expr_values(expr);
328            normalize_expr_values(lower);
329            normalize_expr_values(upper);
330        }
331        Expr::IsNull(expr) | Expr::IsNotNull(expr) | Expr::Not(expr) => {
332            normalize_expr_values(expr);
333        }
334        Expr::Column(_) => {}
335    }
336}
337
338fn collect_select_params(entity: &EntityDescriptor, query: &SelectQuery) -> Vec<Value> {
339    let mut params = Vec::new();
340    if query.raw_sql.is_some() {
341        return params;
342    }
343    for projection in &query.expr_projection {
344        collect_expr_params(&projection.expr, &mut params);
345    }
346    let partitioned = query.partition_by.is_some() && query.slice.is_some();
347    if partitioned {
348        for order in &query.order_by {
349            if let Some(expr) = &order.expr {
350                collect_expr_params(expr, &mut params);
351            }
352        }
353    }
354    if let Some(filter) = &query.filter {
355        collect_expr_params(filter, &mut params);
356    }
357    if let Some(search_text) = &query.search_with_text {
358        let value = Value::from(format!("%{search_text}%"));
359        params.extend(
360            entity
361                .properties
362                .iter()
363                .filter(|property| {
364                    matches!(
365                        property.data_type,
366                        teaql_core::DataType::Text | teaql_core::DataType::LargeText
367                    )
368                })
369                .map(|_| value.clone()),
370        );
371    }
372    if partitioned {
373        return params;
374    }
375    if let Some(having) = &query.having {
376        collect_expr_params(having, &mut params);
377    }
378    for order in &query.order_by {
379        if let Some(expr) = &order.expr {
380            collect_expr_params(expr, &mut params);
381        }
382    }
383    params
384}
385
386fn collect_expr_params(expr: &Expr, params: &mut Vec<Value>) {
387    match expr {
388        Expr::Column(_) => {}
389        Expr::Value(value) => params.push(value.clone()),
390        Expr::Function { args, .. } | Expr::And(args) | Expr::Or(args) => {
391            for arg in args {
392                collect_expr_params(arg, params);
393            }
394        }
395        Expr::Binary { left, op, right } => {
396            collect_expr_params(left, params);
397            if matches!(
398                op,
399                teaql_core::BinaryOp::In
400                    | teaql_core::BinaryOp::NotIn
401                    | teaql_core::BinaryOp::InLarge
402                    | teaql_core::BinaryOp::NotInLarge
403            ) && let Expr::Value(Value::List(values)) = right.as_ref()
404            {
405                params.extend(values.iter().cloned());
406            } else {
407                collect_expr_params(right, params);
408            }
409        }
410        Expr::SubQuery {
411            left,
412            entity,
413            query,
414            ..
415        } => {
416            collect_expr_params(left, params);
417            params.extend(collect_select_params(entity, query));
418        }
419        Expr::Between { expr, lower, upper } => {
420            collect_expr_params(expr, params);
421            collect_expr_params(lower, params);
422            collect_expr_params(upper, params);
423        }
424        Expr::IsNull(expr) | Expr::IsNotNull(expr) | Expr::Not(expr) => {
425            collect_expr_params(expr, params);
426        }
427    }
428}
429
430impl<D, T, S> SqlDataServiceExecutor<D, T, S>
431where
432    S: teaql_data_service::SchemaProvider,
433{
434    fn entity_descriptor(&self, name: &str) -> Option<Arc<teaql_core::EntityDescriptor>> {
435        if let Ok(cache) = self.descriptor_cache.read() {
436            if let Some(descriptor) = cache.get(name) {
437                return Some(descriptor.clone());
438            }
439        }
440        let descriptor = self.schema_provider.get_entity(name)?;
441        if let Ok(mut cache) = self.descriptor_cache.write() {
442            return Some(
443                cache
444                    .entry(name.to_owned())
445                    .or_insert_with(|| descriptor.clone())
446                    .clone(),
447            );
448        }
449        Some(descriptor)
450    }
451}
452
453impl<
454    D: SqlDialect + Send + Sync,
455    T: SqlTransport + Send + Sync,
456    S: teaql_data_service::SchemaProvider + Send + Sync,
457> DataServiceExecutor for SqlDataServiceExecutor<D, T, S>
458{
459    type Error = SqlExecutorError<T::Error>;
460
461    fn capabilities(&self) -> DataServiceCapabilities {
462        DataServiceCapabilities {
463            query: true,
464            mutation: true,
465            transaction: false, // Override if T implements SqlTransactionTransport
466            schema: false,
467            id_generation: false,
468            batch_mutation: true,
469            returning: false,
470        }
471    }
472}
473
474#[cfg(test)]
475mod tests {
476    use super::*;
477    use std::sync::atomic::{AtomicUsize, Ordering};
478    use teaql_core::{DataType, EntityDescriptor, PropertyDescriptor};
479
480    #[derive(Clone, Copy)]
481    struct TestDialect;
482
483    impl SqlDialect for TestDialect {
484        fn kind(&self) -> crate::DatabaseKind {
485            crate::DatabaseKind::PostgreSql
486        }
487
488        fn quote_ident(&self, ident: &str) -> String {
489            format!("\"{ident}\"")
490        }
491
492        fn placeholder(&self, index: usize) -> String {
493            format!("${index}")
494        }
495    }
496
497    #[derive(Clone, Copy)]
498    struct EmptyTransport;
499
500    impl SqlTransport for EmptyTransport {
501        type Error = std::io::Error;
502
503        async fn fetch_all_sql(&self, _query: &CompiledQuery) -> Result<Vec<Record>, Self::Error> {
504            Ok(Vec::new())
505        }
506
507        async fn execute_sql(&self, _query: &CompiledQuery) -> Result<u64, Self::Error> {
508            Ok(0)
509        }
510    }
511
512    #[derive(Clone)]
513    struct CountingSchemaProvider {
514        lookups: Arc<AtomicUsize>,
515    }
516
517    impl teaql_data_service::SchemaProvider for CountingSchemaProvider {
518        fn get_entity(&self, name: &str) -> Option<Arc<EntityDescriptor>> {
519            self.lookups.fetch_add(1, Ordering::Relaxed);
520            (name == "Order").then(|| Arc::new(test_entity()))
521        }
522    }
523
524    fn test_entity() -> EntityDescriptor {
525        EntityDescriptor::new("Order")
526            .property(PropertyDescriptor::new("id", DataType::U64).id().not_null())
527            .property(PropertyDescriptor::new("name", DataType::Text))
528    }
529
530    fn query_request(capture_debug_query: bool) -> QueryRequest {
531        QueryRequest {
532            query: SelectQuery::new("Order"),
533            trace_chain: Vec::new(),
534            comment: None,
535            capture_debug_query,
536        }
537    }
538
539    #[tokio::test]
540    async fn caches_entity_descriptors_across_executor_clones() {
541        let lookups = Arc::new(AtomicUsize::new(0));
542        let executor = SqlDataServiceExecutor::new(
543            TestDialect,
544            EmptyTransport,
545            CountingSchemaProvider {
546                lookups: lookups.clone(),
547            },
548        );
549
550        let result = executor.query(query_request(false)).await.unwrap();
551        executor.clone().query(query_request(true)).await.unwrap();
552
553        assert_eq!(lookups.load(Ordering::Relaxed), 1);
554        assert!(result.metadata.debug_query.is_none());
555    }
556
557    #[tokio::test]
558    async fn cached_select_plan_rebinds_values_and_separates_in_list_lengths() {
559        let lookups = Arc::new(AtomicUsize::new(0));
560        let executor = SqlDataServiceExecutor::new(
561            TestDialect,
562            EmptyTransport,
563            CountingSchemaProvider { lookups },
564        );
565        let request = |filter| QueryRequest {
566            query: SelectQuery::new("Order").filter(filter),
567            trace_chain: Vec::new(),
568            comment: None,
569            capture_debug_query: false,
570        };
571
572        let first = executor
573            .query(request(Expr::eq("id", 7_u64)))
574            .await
575            .unwrap();
576        let second = executor
577            .query(request(Expr::eq("id", 9_u64)))
578            .await
579            .unwrap();
580        assert_eq!(
581            first.metadata.parameterized_query,
582            second.metadata.parameterized_query
583        );
584        assert_eq!(first.metadata.params, vec![Value::U64(7)]);
585        assert_eq!(second.metadata.params, vec![Value::U64(9)]);
586
587        let short = executor
588            .query(request(Expr::in_list("id", [Value::U64(1), Value::U64(2)])))
589            .await
590            .unwrap();
591        let long = executor
592            .query(request(Expr::in_list(
593                "id",
594                [Value::U64(1), Value::U64(2), Value::U64(3)],
595            )))
596            .await
597            .unwrap();
598        assert_ne!(
599            short.metadata.parameterized_query,
600            long.metadata.parameterized_query
601        );
602        assert_eq!(short.metadata.params.len(), 2);
603        assert_eq!(long.metadata.params.len(), 3);
604    }
605
606    #[tokio::test]
607    async fn cached_select_plan_preserves_parameter_order_for_supported_query_shapes() {
608        let executor = SqlDataServiceExecutor::new(
609            TestDialect,
610            EmptyTransport,
611            CountingSchemaProvider {
612                lookups: Arc::new(AtomicUsize::new(0)),
613            },
614        );
615
616        async fn assert_rebound(
617            executor: &SqlDataServiceExecutor<TestDialect, EmptyTransport, CountingSchemaProvider>,
618            warm: SelectQuery,
619            current: SelectQuery,
620        ) {
621            let request = |query| QueryRequest {
622                query,
623                trace_chain: Vec::new(),
624                comment: None,
625                capture_debug_query: false,
626            };
627            executor.query(request(warm)).await.unwrap();
628            let actual = executor.query(request(current.clone())).await.unwrap();
629            let expected = TestDialect
630                .compile_select(&test_entity(), &current)
631                .unwrap();
632            assert_eq!(actual.metadata.parameterized_query, Some(expected.sql));
633            assert_eq!(actual.metadata.params, expected.params);
634        }
635
636        assert_rebound(
637            &executor,
638            SelectQuery::new("Order").search_with_text("first"),
639            SelectQuery::new("Order").search_with_text("second"),
640        )
641        .await;
642        assert_rebound(
643            &executor,
644            SelectQuery::new("Order")
645                .project_expr("marker", Expr::value(1_i64))
646                .filter(Expr::eq("id", 2_u64))
647                .having(Expr::gt("id", 3_u64))
648                .order_by(teaql_core::OrderBy::asc_expr(Expr::value(4_i64))),
649            SelectQuery::new("Order")
650                .project_expr("marker", Expr::value(11_i64))
651                .filter(Expr::eq("id", 12_u64))
652                .having(Expr::gt("id", 13_u64))
653                .order_by(teaql_core::OrderBy::asc_expr(Expr::value(14_i64))),
654        )
655        .await;
656        assert_rebound(
657            &executor,
658            SelectQuery::new("Order")
659                .filter(Expr::eq("id", 1_u64))
660                .order_by(teaql_core::OrderBy::asc_expr(Expr::value(2_i64)))
661                .page(0, 10)
662                .partition_by("name"),
663            SelectQuery::new("Order")
664                .filter(Expr::eq("id", 3_u64))
665                .order_by(teaql_core::OrderBy::asc_expr(Expr::value(4_i64)))
666                .page(0, 10)
667                .partition_by("name"),
668        )
669        .await;
670        assert_rebound(
671            &executor,
672            SelectQuery::new("Order").filter(Expr::in_subquery(
673                "id",
674                test_entity(),
675                SelectQuery::new("Order").filter(Expr::gt("id", 20_u64)),
676                "id",
677            )),
678            SelectQuery::new("Order").filter(Expr::in_subquery(
679                "id",
680                test_entity(),
681                SelectQuery::new("Order").filter(Expr::gt("id", 30_u64)),
682                "id",
683            )),
684        )
685        .await;
686    }
687}
688
689impl<
690    D: SqlDialect + Send + Sync,
691    T: SqlTransport + Send + Sync,
692    S: teaql_data_service::SchemaProvider + Send + Sync,
693> QueryExecutor for SqlDataServiceExecutor<D, T, S>
694{
695    fn query(
696        &self,
697        request: QueryRequest,
698    ) -> impl std::future::Future<Output = Result<QueryResult, Self::Error>> + Send {
699        async move {
700            let entity_desc = self
701                .entity_descriptor(&request.query.entity)
702                .ok_or_else(|| {
703                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(
704                        request.query.entity.clone(),
705                    ))
706                })?;
707
708            let compiled = self
709                .compile_select_cached(&entity_desc, &request.query)
710                .map_err(SqlExecutorError::Compile)?;
711            let start = SystemTime::now();
712            let rows = self
713                .transport
714                .fetch_all_sql(&compiled)
715                .await
716                .map_err(SqlExecutorError::Transport)?;
717            let end = SystemTime::now();
718            let debug_query = request
719                .capture_debug_query
720                .then(|| compiled.debug_sql(self.dialect.kind()));
721            let CompiledQuery { sql, params, .. } = compiled;
722
723            let metadata = ExecutionMetadata {
724                backend: "sql".to_string(),
725                operation: DataServiceOperation::Query,
726                started_at: start,
727                ended_at: end,
728                affected_rows: None,
729                result_count: Some(rows.len()),
730                trace_chain: request.trace_chain,
731                comment: request.comment,
732                backend_request_id: None,
733                parameterized_query: Some(sql),
734                params,
735                debug_query,
736            };
737
738            Ok(QueryResult { rows, metadata })
739        }
740    }
741}
742
743impl<
744    D: SqlDialect + Send + Sync,
745    T: SqlTransport + Send + Sync,
746    S: teaql_data_service::SchemaProvider + Send + Sync,
747> MutationExecutor for SqlDataServiceExecutor<D, T, S>
748{
749    fn mutate(
750        &self,
751        request: MutationRequest,
752    ) -> impl std::future::Future<Output = Result<MutationResult, Self::Error>> + Send {
753        async move {
754            let entity_name = match &request {
755                MutationRequest::Insert(cmd) => &cmd.entity,
756                MutationRequest::Update(cmd) => &cmd.entity,
757                MutationRequest::Delete(cmd) => &cmd.entity,
758                MutationRequest::Recover(cmd) => &cmd.entity,
759                MutationRequest::Batch(mutations) => {
760                    let mut total_affected = 0;
761                    let mut parameterized_queries = Vec::new();
762                    let mut params = Vec::new();
763                    let mut debug_queries = Vec::new();
764                    let start = SystemTime::now();
765                    for req in mutations {
766                        let res = Box::pin(self.mutate(req.clone())).await?;
767                        total_affected += res.affected_rows;
768                        if let Some(query) = res.metadata.parameterized_query {
769                            parameterized_queries.push(query);
770                        }
771                        params.extend(res.metadata.params);
772                        if let Some(query) = res.metadata.debug_query {
773                            debug_queries.push(query);
774                        }
775                    }
776                    let end = SystemTime::now();
777                    return Ok(MutationResult {
778                        affected_rows: total_affected,
779                        generated_values: Record::default(),
780                        persisted_record: None,
781                        metadata: ExecutionMetadata {
782                            backend: "sql".to_string(),
783                            operation: DataServiceOperation::Batch,
784                            started_at: start,
785                            ended_at: end,
786                            affected_rows: Some(total_affected),
787                            result_count: None,
788                            trace_chain: Vec::new(),
789                            comment: None,
790                            backend_request_id: None,
791                            parameterized_query: (!parameterized_queries.is_empty())
792                                .then(|| parameterized_queries.join("; ")),
793                            params,
794                            debug_query: (!debug_queries.is_empty())
795                                .then(|| debug_queries.join("; ")),
796                        },
797                    });
798                }
799            };
800
801            let entity_desc = self.entity_descriptor(entity_name).ok_or_else(|| {
802                SqlExecutorError::Compile(SqlCompileError::UnknownEntity(entity_name.clone()))
803            })?;
804
805            let compiled = match &request {
806                MutationRequest::Insert(cmd) => self
807                    .dialect
808                    .compile_insert(&entity_desc, cmd)
809                    .map_err(SqlExecutorError::Compile)?,
810                MutationRequest::Update(cmd) => self
811                    .dialect
812                    .compile_update(&entity_desc, cmd)
813                    .map_err(SqlExecutorError::Compile)?,
814                MutationRequest::Delete(cmd) => self
815                    .dialect
816                    .compile_delete(&entity_desc, cmd)
817                    .map_err(SqlExecutorError::Compile)?,
818                MutationRequest::Recover(cmd) => self
819                    .dialect
820                    .compile_recover(&entity_desc, cmd)
821                    .map_err(SqlExecutorError::Compile)?,
822                MutationRequest::Batch(_) => unreachable!(),
823            };
824
825            let start = SystemTime::now();
826            let affected_rows = self
827                .transport
828                .execute_sql(&compiled)
829                .await
830                .map_err(SqlExecutorError::Transport)?;
831            let end = SystemTime::now();
832
833            let operation = match &request {
834                MutationRequest::Insert(_) => DataServiceOperation::Insert,
835                MutationRequest::Update(_) => DataServiceOperation::Update,
836                MutationRequest::Delete(_) => DataServiceOperation::Delete,
837                MutationRequest::Recover(_) => DataServiceOperation::Recover,
838                MutationRequest::Batch(_) => DataServiceOperation::Batch,
839            };
840
841            let metadata = ExecutionMetadata {
842                backend: "sql".to_string(),
843                operation,
844                started_at: start,
845                ended_at: end,
846                affected_rows: Some(affected_rows),
847                result_count: None,
848                trace_chain: request.trace_chain().to_vec(),
849                comment: request.comment().map(|s| s.to_owned()),
850                backend_request_id: None,
851                parameterized_query: Some(compiled.sql.clone()),
852                params: compiled.params.clone(),
853                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
854            };
855
856            Ok(MutationResult {
857                affected_rows,
858                generated_values: Record::default(),
859                persisted_record: None,
860                metadata,
861            })
862        }
863    }
864}
865
866#[derive(Clone)]
867pub struct SqlDataServiceTransaction<'a, D, Tx: SqlTransport + SqlTransaction, S> {
868    pub dialect: &'a D,
869    pub transport: Tx,
870    pub schema_provider: &'a S,
871    descriptor_cache: Arc<RwLock<HashMap<String, Arc<teaql_core::EntityDescriptor>>>>,
872    select_plan_cache: Arc<RwLock<Vec<(SelectQuery, String)>>>,
873}
874
875impl<'a, D, Tx: SqlTransport + SqlTransaction, S> SqlDataServiceTransaction<'a, D, Tx, S>
876where
877    S: teaql_data_service::SchemaProvider,
878{
879    fn entity_descriptor(&self, name: &str) -> Option<Arc<teaql_core::EntityDescriptor>> {
880        if let Ok(cache) = self.descriptor_cache.read() {
881            if let Some(descriptor) = cache.get(name) {
882                return Some(descriptor.clone());
883            }
884        }
885        let descriptor = self.schema_provider.get_entity(name)?;
886        if let Ok(mut cache) = self.descriptor_cache.write() {
887            return Some(
888                cache
889                    .entry(name.to_owned())
890                    .or_insert_with(|| descriptor.clone())
891                    .clone(),
892            );
893        }
894        Some(descriptor)
895    }
896
897    fn compile_select_cached(
898        &self,
899        entity: &EntityDescriptor,
900        query: &SelectQuery,
901    ) -> Result<CompiledQuery, SqlCompileError>
902    where
903        D: SqlDialect,
904    {
905        compile_select_with_cache(self.dialect, &self.select_plan_cache, entity, query)
906    }
907}
908
909impl<
910    'a,
911    D: SqlDialect + Send + Sync,
912    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
913    S: teaql_data_service::SchemaProvider + Send + Sync,
914> DataServiceExecutor for SqlDataServiceTransaction<'a, D, Tx, S>
915{
916    type Error = SqlExecutorError<<Tx as SqlTransport>::Error>;
917
918    fn capabilities(&self) -> DataServiceCapabilities {
919        DataServiceCapabilities {
920            query: true,
921            mutation: true,
922            transaction: false,
923            schema: false,
924            id_generation: false,
925            batch_mutation: true,
926            returning: false,
927        }
928    }
929}
930
931impl<
932    'a,
933    D: SqlDialect + Send + Sync,
934    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
935    S: teaql_data_service::SchemaProvider + Send + Sync,
936> QueryExecutor for SqlDataServiceTransaction<'a, D, Tx, S>
937{
938    fn query(
939        &self,
940        request: QueryRequest,
941    ) -> impl std::future::Future<Output = Result<QueryResult, Self::Error>> + Send {
942        async move {
943            let entity_desc = self
944                .entity_descriptor(&request.query.entity)
945                .ok_or_else(|| {
946                    SqlExecutorError::Compile(SqlCompileError::UnknownEntity(
947                        request.query.entity.clone(),
948                    ))
949                })?;
950
951            let compiled = self
952                .compile_select_cached(&entity_desc, &request.query)
953                .map_err(SqlExecutorError::Compile)?;
954            let start = SystemTime::now();
955            let rows = self
956                .transport
957                .fetch_all_sql(&compiled)
958                .await
959                .map_err(SqlExecutorError::Transport)?;
960            let end = SystemTime::now();
961
962            let metadata = ExecutionMetadata {
963                backend: "sql".to_string(),
964                operation: DataServiceOperation::Query,
965                started_at: start,
966                ended_at: end,
967                affected_rows: None,
968                result_count: Some(rows.len()),
969                trace_chain: request.trace_chain,
970                comment: request.comment,
971                backend_request_id: None,
972                parameterized_query: Some(compiled.sql.clone()),
973                params: compiled.params.clone(),
974                debug_query: request
975                    .capture_debug_query
976                    .then(|| compiled.debug_sql(self.dialect.kind())),
977            };
978
979            Ok(QueryResult { rows, metadata })
980        }
981    }
982}
983
984impl<
985    'a,
986    D: SqlDialect + Send + Sync,
987    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
988    S: teaql_data_service::SchemaProvider + Send + Sync,
989> MutationExecutor for SqlDataServiceTransaction<'a, D, Tx, S>
990{
991    fn mutate(
992        &self,
993        request: MutationRequest,
994    ) -> impl std::future::Future<Output = Result<MutationResult, Self::Error>> + Send {
995        async move {
996            let entity_name = match &request {
997                MutationRequest::Insert(cmd) => &cmd.entity,
998                MutationRequest::Update(cmd) => &cmd.entity,
999                MutationRequest::Delete(cmd) => &cmd.entity,
1000                MutationRequest::Recover(cmd) => &cmd.entity,
1001                MutationRequest::Batch(mutations) => {
1002                    let mut total_affected = 0;
1003                    let mut parameterized_queries = Vec::new();
1004                    let mut params = Vec::new();
1005                    let mut debug_queries = Vec::new();
1006                    let start = SystemTime::now();
1007                    for req in mutations {
1008                        let res = Box::pin(self.mutate(req.clone())).await?;
1009                        total_affected += res.affected_rows;
1010                        if let Some(query) = res.metadata.parameterized_query {
1011                            parameterized_queries.push(query);
1012                        }
1013                        params.extend(res.metadata.params);
1014                        if let Some(query) = res.metadata.debug_query {
1015                            debug_queries.push(query);
1016                        }
1017                    }
1018                    let end = SystemTime::now();
1019                    return Ok(MutationResult {
1020                        affected_rows: total_affected,
1021                        generated_values: Record::default(),
1022                        persisted_record: None,
1023                        metadata: ExecutionMetadata {
1024                            backend: "sql".to_string(),
1025                            operation: DataServiceOperation::Batch,
1026                            started_at: start,
1027                            ended_at: end,
1028                            affected_rows: Some(total_affected),
1029                            result_count: None,
1030                            trace_chain: Vec::new(),
1031                            comment: None,
1032                            backend_request_id: None,
1033                            parameterized_query: (!parameterized_queries.is_empty())
1034                                .then(|| parameterized_queries.join("; ")),
1035                            params,
1036                            debug_query: (!debug_queries.is_empty())
1037                                .then(|| debug_queries.join("; ")),
1038                        },
1039                    });
1040                }
1041            };
1042
1043            let entity_desc = self.entity_descriptor(entity_name).ok_or_else(|| {
1044                SqlExecutorError::Compile(SqlCompileError::UnknownEntity(entity_name.clone()))
1045            })?;
1046
1047            let compiled = match &request {
1048                MutationRequest::Insert(cmd) => self
1049                    .dialect
1050                    .compile_insert(&entity_desc, cmd)
1051                    .map_err(SqlExecutorError::Compile)?,
1052                MutationRequest::Update(cmd) => self
1053                    .dialect
1054                    .compile_update(&entity_desc, cmd)
1055                    .map_err(SqlExecutorError::Compile)?,
1056                MutationRequest::Delete(cmd) => self
1057                    .dialect
1058                    .compile_delete(&entity_desc, cmd)
1059                    .map_err(SqlExecutorError::Compile)?,
1060                MutationRequest::Recover(cmd) => self
1061                    .dialect
1062                    .compile_recover(&entity_desc, cmd)
1063                    .map_err(SqlExecutorError::Compile)?,
1064                MutationRequest::Batch(_) => unreachable!("batch handled above"),
1065            };
1066
1067            let start = SystemTime::now();
1068            let affected_rows = self
1069                .transport
1070                .execute_sql(&compiled)
1071                .await
1072                .map_err(SqlExecutorError::Transport)?;
1073            let end = SystemTime::now();
1074
1075            let operation = match &request {
1076                MutationRequest::Insert(_) => DataServiceOperation::Insert,
1077                MutationRequest::Update(_) => DataServiceOperation::Update,
1078                MutationRequest::Delete(_) => DataServiceOperation::Delete,
1079                MutationRequest::Recover(_) => DataServiceOperation::Recover,
1080                MutationRequest::Batch(_) => DataServiceOperation::Batch,
1081            };
1082
1083            let persisted_id = match &request {
1084                MutationRequest::Insert(cmd) => cmd.values.get("id").cloned(),
1085                MutationRequest::Update(cmd) => Some(cmd.id.clone()),
1086                MutationRequest::Delete(cmd) if cmd.soft_delete => Some(cmd.id.clone()),
1087                MutationRequest::Recover(cmd) => Some(cmd.id.clone()),
1088                MutationRequest::Delete(_) | MutationRequest::Batch(_) => None,
1089            };
1090            let persisted_record = if affected_rows == 1 {
1091                if let Some(id) = persisted_id {
1092                    let query = SelectQuery::new(entity_name.clone()).filter(Expr::eq("id", id));
1093                    let compiled_readback = self
1094                        .compile_select_cached(&entity_desc, &query)
1095                        .map_err(SqlExecutorError::Compile)?;
1096                    let mut rows = self
1097                        .transport
1098                        .fetch_all_sql(&compiled_readback)
1099                        .await
1100                        .map_err(SqlExecutorError::Transport)?;
1101                    if rows.len() != 1 {
1102                        return Err(SqlExecutorError::PersistedRecord(format!(
1103                            "persisted {entity_name} record could not be read back"
1104                        )));
1105                    }
1106                    rows.pop()
1107                } else {
1108                    None
1109                }
1110            } else {
1111                None
1112            };
1113
1114            let metadata = ExecutionMetadata {
1115                backend: "sql".to_string(),
1116                operation,
1117                started_at: start,
1118                ended_at: end,
1119                affected_rows: Some(affected_rows),
1120                result_count: None,
1121                trace_chain: request.trace_chain().to_vec(),
1122                comment: request.comment().map(|s| s.to_owned()),
1123                backend_request_id: None,
1124                parameterized_query: Some(compiled.sql.clone()),
1125                params: compiled.params.clone(),
1126                debug_query: Some(compiled.debug_sql(self.dialect.kind())),
1127            };
1128
1129            Ok(MutationResult {
1130                affected_rows,
1131                generated_values: Record::default(),
1132                persisted_record,
1133                metadata,
1134            })
1135        }
1136    }
1137}
1138
1139impl<
1140    'a,
1141    D: SqlDialect + Send + Sync,
1142    Tx: SqlTransport + SqlTransaction<Error = <Tx as SqlTransport>::Error> + Send + Sync,
1143    S: teaql_data_service::SchemaProvider + Send + Sync,
1144> teaql_data_service::Transaction for SqlDataServiceTransaction<'a, D, Tx, S>
1145{
1146    type Error = SqlExecutorError<<Tx as SqlTransport>::Error>;
1147
1148    fn commit(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send {
1149        async move {
1150            self.transport
1151                .commit_sql()
1152                .await
1153                .map_err(SqlExecutorError::Transport)
1154        }
1155    }
1156
1157    fn rollback(self) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send {
1158        async move {
1159            self.transport
1160                .rollback_sql()
1161                .await
1162                .map_err(SqlExecutorError::Transport)
1163        }
1164    }
1165}
1166
1167impl<
1168    D: SqlDialect + Send + Sync,
1169    T: SqlTransactionTransport + Send + Sync,
1170    S: teaql_data_service::SchemaProvider + Send + Sync,
1171> teaql_data_service::TransactionExecutor for SqlDataServiceExecutor<D, T, S>
1172{
1173    type Tx<'a>
1174        = SqlDataServiceTransaction<'a, D, T::Tx<'a>, S>
1175    where
1176        Self: 'a;
1177
1178    fn begin(&self) -> impl std::future::Future<Output = Result<Self::Tx<'_>, Self::Error>> + Send {
1179        async move {
1180            let tx = self
1181                .transport
1182                .begin_sql()
1183                .await
1184                .map_err(SqlExecutorError::Transport)?;
1185            Ok(SqlDataServiceTransaction {
1186                dialect: &self.dialect,
1187                transport: tx,
1188                schema_provider: &self.schema_provider,
1189                descriptor_cache: self.descriptor_cache.clone(),
1190                select_plan_cache: self.select_plan_cache.clone(),
1191            })
1192        }
1193    }
1194}
1195
1196impl<
1197    D: SqlDialect + Send + Sync,
1198    T: StreamingSqlTransport + Send + Sync,
1199    S: teaql_data_service::SchemaProvider + Send + Sync,
1200> teaql_data_service::StreamQueryExecutor for SqlDataServiceExecutor<D, T, S>
1201{
1202    fn query_stream(
1203        &self,
1204        request: teaql_data_service::QueryRequest,
1205        chunk_size: usize,
1206    ) -> teaql_data_service::QueryStream<'_, Self::Error> {
1207        use futures_util::StreamExt;
1208        let entity = match self.entity_descriptor(&request.query.entity) {
1209            Some(entity) => entity,
1210            None => {
1211                return Box::pin(futures_util::stream::once(async {
1212                    Err(SqlExecutorError::Compile(SqlCompileError::UnknownEntity(
1213                        request.query.entity,
1214                    )))
1215                }));
1216            }
1217        };
1218        match self.compile_select_cached(&entity, &request.query) {
1219            Ok(compiled) => Box::pin(
1220                self.transport
1221                    .stream_sql(compiled, chunk_size)
1222                    .map(|r| r.map_err(SqlExecutorError::Transport)),
1223            ),
1224            Err(error) => Box::pin(futures_util::stream::once(async {
1225                Err(SqlExecutorError::Compile(error))
1226            })),
1227        }
1228    }
1229}