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