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, 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(), ¤t)
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}