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