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