1#![allow(warnings)]
2use std::collections::BTreeMap;
3use std::future::Future;
4use std::pin::Pin;
5
6use chrono::{DateTime, NaiveDate, NaiveDateTime, Utc};
7use deadpool_postgres::Pool;
8use rust_decimal::Decimal;
9use std::sync::Arc;
10use teaql_core::{
11 BinaryOp, DataType, EntityDescriptor, Expr, InsertCommand, PropertyDescriptor, Record,
12 SelectQuery, UpdateCommand, Value,
13};
14use teaql_runtime::{GraphNode, InternalIdGenerator, RuntimeError, SchemaProvider, UserContext};
15use teaql_sql::{
16 CompiledQuery, DatabaseKind, SqlCompileError, SqlDialect, SqlTransport,
17 quote_identifier_if_needed,
18};
19use tokio::sync::Mutex;
20
21pub const DEFAULT_ID_SPACE_TABLE: &str = "teaql_id_space";
22
23#[derive(Debug, Default, Clone, Copy)]
24pub struct PostgresDialect;
25
26impl SqlDialect for PostgresDialect {
27 fn kind(&self) -> DatabaseKind {
28 DatabaseKind::PostgreSql
29 }
30
31 fn quote_ident(&self, ident: &str) -> String {
32 quote_ident(ident)
33 }
34
35 fn placeholder(&self, index: usize) -> String {
36 format!("${index}")
37 }
38
39 fn schema_setup_sqls(&self) -> &'static [&'static str] {
40 &[CREATE_SOUNDEX_FUNCTION]
41 }
42
43 fn schema_type_sql(
44 &self,
45 data_type: DataType,
46 _property: &PropertyDescriptor,
47 ) -> Result<&'static str, SqlCompileError> {
48 match data_type {
49 DataType::Bool => Ok("BOOLEAN"),
50 DataType::I64 | DataType::U64 => Ok("BIGINT"),
51 DataType::F64 => Ok("DOUBLE PRECISION"),
52 DataType::Decimal => Ok("NUMERIC"),
53 DataType::Text => Ok("VARCHAR(255)"),
54 DataType::LargeText => Ok("TEXT"),
55 DataType::Json => Ok("JSONB"),
56 DataType::Date => Ok("DATE"),
57 DataType::Timestamp => Ok("TIMESTAMPTZ"),
58 }
59 }
60
61 fn compile_in(
62 &self,
63 entity: &EntityDescriptor,
64 left: &Expr,
65 op: BinaryOp,
66 right: &Expr,
67 params: &mut Vec<Value>,
68 ) -> Result<String, SqlCompileError> {
69 match op {
70 BinaryOp::InLarge | BinaryOp::NotInLarge => {
71 let Expr::Value(Value::List(values)) = right else {
72 let lhs = self.compile_expr(entity, left, params)?;
73 let rhs = self.compile_expr(entity, right, params)?;
74 let operator = match op {
75 BinaryOp::InLarge => "= ANY",
76 BinaryOp::NotInLarge => "<> ALL",
77 _ => unreachable!(),
78 };
79 return Ok(format!("({lhs} {operator} ({rhs}))"));
80 };
81 if values.is_empty() {
82 return Err(SqlCompileError::EmptyInList);
83 }
84 let lhs = self.compile_expr(entity, left, params)?;
85 params.push(Value::List(values.clone()));
86 let placeholder = self.placeholder(params.len());
87 let operator = match op {
88 BinaryOp::InLarge => "= ANY",
89 BinaryOp::NotInLarge => "<> ALL",
90 _ => unreachable!(),
91 };
92 Ok(format!("({lhs} {operator}({placeholder}))"))
93 }
94 _ => {
95 let lhs = self.compile_expr(entity, left, params)?;
96 let operator = match op {
97 BinaryOp::In => "IN",
98 BinaryOp::NotIn => "NOT IN",
99 _ => unreachable!(),
100 };
101 match right {
102 Expr::Value(Value::List(values)) => {
103 if values.is_empty() {
104 return Err(SqlCompileError::EmptyInList);
105 }
106 let mut placeholders = Vec::with_capacity(values.len());
107 for value in values {
108 params.push(value.clone());
109 placeholders.push(self.placeholder(params.len()));
110 }
111 Ok(format!("({lhs} {operator} ({}))", placeholders.join(", ")))
112 }
113 _ => {
114 let rhs = self.compile_expr(entity, right, params)?;
115 Ok(format!("({lhs} {operator} ({rhs}))"))
116 }
117 }
118 }
119 }
120 }
121}
122
123const CREATE_SOUNDEX_FUNCTION: &str = r#"
124CREATE OR REPLACE FUNCTION soundex(input text)
125RETURNS text
126LANGUAGE plpgsql
127IMMUTABLE
128STRICT
129AS $$
130DECLARE
131 normalized text := upper(regexp_replace(input, '[^A-Za-z]', '', 'g'));
132 first_char text;
133 output text;
134 previous_code text;
135 code text;
136 ch text;
137 i integer;
138BEGIN
139 IF normalized = '' THEN
140 RETURN '0000';
141 END IF;
142
143 first_char := substr(normalized, 1, 1);
144 output := first_char;
145 previous_code := CASE
146 WHEN first_char IN ('B', 'F', 'P', 'V') THEN '1'
147 WHEN first_char IN ('C', 'G', 'J', 'K', 'Q', 'S', 'X', 'Z') THEN '2'
148 WHEN first_char IN ('D', 'T') THEN '3'
149 WHEN first_char = 'L' THEN '4'
150 WHEN first_char IN ('M', 'N') THEN '5'
151 WHEN first_char = 'R' THEN '6'
152 ELSE '0'
153 END;
154
155 FOR i IN 2..char_length(normalized) LOOP
156 ch := substr(normalized, i, 1);
157 code := CASE
158 WHEN ch IN ('B', 'F', 'P', 'V') THEN '1'
159 WHEN ch IN ('C', 'G', 'J', 'K', 'Q', 'S', 'X', 'Z') THEN '2'
160 WHEN ch IN ('D', 'T') THEN '3'
161 WHEN ch = 'L' THEN '4'
162 WHEN ch IN ('M', 'N') THEN '5'
163 WHEN ch = 'R' THEN '6'
164 ELSE '0'
165 END;
166
167 IF code <> '0' AND code <> previous_code THEN
168 output := output || code;
169 IF char_length(output) = 4 THEN
170 RETURN output;
171 END IF;
172 END IF;
173 previous_code := code;
174 END LOOP;
175
176 RETURN rpad(output, 4, '0');
177END;
178$$
179"#;
180
181#[derive(Debug)]
182pub enum MutationExecutorError {
183 Driver(tokio_postgres::Error),
184 Pool(String),
185 SqlCompile(SqlCompileError),
186 UnsupportedValue(&'static str),
187 UnsupportedColumnType(String),
188 Bind(String),
189}
190
191impl std::fmt::Display for MutationExecutorError {
192 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
193 match self {
194 Self::Driver(err) => err.fmt(f),
195 Self::Pool(err) => write!(f, "postgres pool error: {err}"),
196 Self::SqlCompile(err) => err.fmt(f),
197 Self::UnsupportedValue(kind) => {
198 write!(f, "unsupported bind value for mutation executor: {kind}")
199 }
200 Self::UnsupportedColumnType(kind) => {
201 write!(f, "unsupported column type for record decoding: {kind}")
202 }
203 Self::Bind(message) => write!(f, "bind error: {message}"),
204 }
205 }
206}
207
208impl std::error::Error for MutationExecutorError {}
209
210impl From<tokio_postgres::Error> for MutationExecutorError {
211 fn from(value: tokio_postgres::Error) -> Self {
212 Self::Driver(value)
213 }
214}
215
216impl From<SqlCompileError> for MutationExecutorError {
217 fn from(value: SqlCompileError) -> Self {
218 Self::SqlCompile(value)
219 }
220}
221
222#[derive(Clone)]
223pub struct PgMutationExecutor {
224 pool: Pool,
225}
226
227impl SqlTransport for PgMutationExecutor {
228 type Error = MutationExecutorError;
229
230 async fn fetch_all_sql(&self, query: &CompiledQuery) -> Result<Vec<Record>, Self::Error> {
231 self.fetch_all(query).await
232 }
233
234 async fn execute_sql(&self, query: &CompiledQuery) -> Result<u64, Self::Error> {
235 self.execute(query).await
236 }
237}
238
239impl teaql_sql::StreamingSqlTransport for PgMutationExecutor {
240 fn stream_sql(
241 &self,
242 query: CompiledQuery,
243 chunk_size: usize,
244 ) -> teaql_data_service::QueryStream<'_, Self::Error> {
245 let pool = self.pool.clone();
246 Box::pin(async_stream::try_stream! {
247 use futures_util::TryStreamExt;
248 let mut args = PgArgs { values: Vec::new() }; for value in &query.params { bind_pg(&mut args, value)?; }
249 let client = pool.get().await.map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
250 let params = args.as_refs();
251 let statement = client.prepare_cached(&query.sql).await?;
252 let rows = client.query_raw(&statement, params).await?;
253 futures_util::pin_mut!(rows);
254 let mut chunk = Vec::with_capacity(chunk_size); let mut index = 0;
255 while let Some(row) = rows.try_next().await? { chunk.push(decode_pg_row(&row)?); if chunk.len()==chunk_size { yield teaql_data_service::StreamChunk { rows: std::mem::take(&mut chunk), chunk_index:index, is_last:false }; index+=1; } }
256 if !chunk.is_empty() { yield teaql_data_service::StreamChunk { rows:chunk, chunk_index:index, is_last:true }; }
257 })
258 }
259}
260
261impl teaql_sql::SqlTransaction for PgMutationExecutor {
262 type Error = MutationExecutorError;
263
264 async fn commit_sql(self) -> Result<(), Self::Error> {
265 Err(MutationExecutorError::Bind(
266 "Transactions not supported yet".to_string(),
267 ))
268 }
269
270 async fn rollback_sql(self) -> Result<(), Self::Error> {
271 Err(MutationExecutorError::Bind(
272 "Transactions not supported yet".to_string(),
273 ))
274 }
275}
276
277impl teaql_sql::SqlTransactionTransport for PgMutationExecutor {
278 type Tx<'a>
279 = Self
280 where
281 Self: 'a;
282
283 async fn begin_sql(&self) -> Result<Self::Tx<'_>, Self::Error> {
284 Err(MutationExecutorError::Bind(
285 "Transactions not supported yet".to_string(),
286 ))
287 }
288}
289
290impl PgMutationExecutor {
291 pub fn new(pool: Pool) -> Self {
292 Self { pool }
293 }
294
295 pub fn pool(&self) -> Pool {
296 self.pool.clone()
297 }
298
299 pub async fn ensure_schema(
300 &self,
301 dialect: &PostgresDialect,
302 entities: &[&EntityDescriptor],
303 ) -> Result<(), MutationExecutorError> {
304 let client = self
305 .pool
306 .get()
307 .await
308 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
309 for sql in dialect.schema_setup_sqls() {
310 client.execute(*sql, &[]).await?;
311 }
312 self.ensure_id_space_table(DEFAULT_ID_SPACE_TABLE).await?;
313
314 for entity in entities {
315 if !self.table_exists(&entity.table_name).await? {
316 let sql = dialect.compile_create_table(entity)?;
317 client.execute(&sql, &[]).await?;
318 continue;
319 }
320
321 let existing_columns = self.table_columns(&entity.table_name).await?;
322 for property in &entity.properties {
323 let bare_column = strip_identifier_quotes(&property.column_name).to_lowercase();
324 if existing_columns.contains(&bare_column) {
325 continue;
326 }
327 let sql = dialect.compile_add_column(entity, property)?;
328 client.execute(&sql, &[]).await?;
329 }
330
331 for sql in dialect.schema_indexes_sqls(entity)? {
332 client.execute(&sql, &[]).await?;
333 }
334 }
335 Ok(())
336 }
337
338 pub async fn ensure_id_space_table(
339 &self,
340 table_name: &str,
341 ) -> Result<(), MutationExecutorError> {
342 let sql = format!(
343 "CREATE TABLE IF NOT EXISTS {} (type_name VARCHAR(100) PRIMARY KEY, current_level BIGINT NOT NULL)",
344 quote_ident(table_name)
345 );
346 let client = self
347 .pool
348 .get()
349 .await
350 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
351 client.execute(&sql, &[]).await?;
352 Ok(())
353 }
354
355 pub async fn execute(&self, query: &CompiledQuery) -> Result<u64, MutationExecutorError> {
356 let mut args = PgArgs { values: Vec::new() };
357 for value in &query.params {
358 bind_pg(&mut args, value)?;
359 }
360 let client = self
361 .pool
362 .get()
363 .await
364 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
365 let statement = client.prepare_cached(&query.sql).await?;
366 let result = client.execute(&statement, &args.as_refs()).await?;
367 Ok(result)
368 }
369
370 pub async fn fetch_all(
371 &self,
372 query: &CompiledQuery,
373 ) -> Result<Vec<Record>, MutationExecutorError> {
374 let mut args = PgArgs { values: Vec::new() };
375 for value in &query.params {
376 bind_pg(&mut args, value)?;
377 }
378 let client = self
379 .pool
380 .get()
381 .await
382 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
383 let statement = client.prepare_cached(&query.sql).await?;
384 let rows = client.query(&statement, &args.as_refs()).await?;
385 rows.iter().map(decode_pg_row).collect()
386 }
387
388 async fn table_exists(&self, table_name: &str) -> Result<bool, MutationExecutorError> {
389 let client = self
390 .pool
391 .get()
392 .await
393 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
394 let row = client
395 .query_one(
396 "SELECT COUNT(1)
397 FROM information_schema.tables
398 WHERE table_schema = current_schema()
399 AND table_name = $1",
400 &[&table_name],
401 )
402 .await?;
403 let exists: i64 = row.try_get(0)?;
404 Ok(exists > 0)
405 }
406
407 async fn table_columns(
408 &self,
409 table_name: &str,
410 ) -> Result<std::collections::BTreeSet<String>, MutationExecutorError> {
411 let client = self
412 .pool
413 .get()
414 .await
415 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
416 let rows = client
417 .query(
418 "SELECT column_name
419 FROM information_schema.columns
420 WHERE table_schema = current_schema()
421 AND table_name = $1",
422 &[&table_name],
423 )
424 .await?;
425 let mut columns = std::collections::BTreeSet::new();
426 for row in rows {
427 let name: String = row.try_get("column_name")?;
428 columns.insert(name.to_lowercase());
429 }
430 Ok(columns)
431 }
432}
433
434async fn ensure_initial_graphs_postgres(
435 executor: &PgMutationExecutor,
436 dialect: &PostgresDialect,
437 context: &UserContext,
438) -> Result<(), MutationExecutorError> {
439 for graph in context.initial_graphs() {
440 let entity = context.entity(&graph.entity).ok_or_else(|| {
441 MutationExecutorError::Bind(format!("missing entity: {}", graph.entity))
442 })?;
443 if initial_graph_exists_postgres(executor, dialect, entity, graph).await? {
444 if let Some(query) = compile_initial_graph_update(dialect, entity, graph)? {
445 executor.execute(&query).await?;
446 }
447 continue;
448 }
449 let query = compile_initial_graph_insert(dialect, entity, graph)?;
450 executor.execute(&query).await?;
451 }
452 for graph in context.root_graphs() {
453 let entity = context.entity(&graph.entity).ok_or_else(|| {
454 MutationExecutorError::Bind(format!("missing entity: {}", graph.entity))
455 })?;
456 if initial_graph_exists_postgres(executor, dialect, entity, graph).await? {
457 continue;
458 }
459 let query = compile_initial_graph_insert(dialect, entity, graph)?;
460 executor.execute(&query).await?;
461 }
462 let generator = PgIdSpaceGenerator::from_executor(executor.clone());
463 for graph in context.initial_graphs().iter().chain(context.root_graphs()) {
464 if let Some(id) = graph.values.get("id").and_then(Value::try_u64) {
465 generator.ensure_floor(&graph.entity, id).await?;
466 }
467 }
468 Ok(())
469}
470
471async fn initial_graph_exists_postgres(
472 executor: &PgMutationExecutor,
473 dialect: &PostgresDialect,
474 entity: &EntityDescriptor,
475 graph: &GraphNode,
476) -> Result<bool, MutationExecutorError> {
477 let Some(id) = graph.values.get("id") else {
478 return Ok(false);
479 };
480 let query = dialect.compile_select(
481 entity,
482 &SelectQuery::new(&graph.entity)
483 .project("id")
484 .filter(Expr::eq("id", id.clone()))
485 .limit(1),
486 )?;
487 Ok(!executor.fetch_all(&query).await?.is_empty())
488}
489
490fn compile_initial_graph_insert(
491 dialect: &impl SqlDialect,
492 entity: &EntityDescriptor,
493 graph: &GraphNode,
494) -> Result<CompiledQuery, MutationExecutorError> {
495 let mut command = InsertCommand::new(&graph.entity);
496 for (field, value) in &graph.values {
497 command = command.value(field.clone(), value.clone());
498 }
499 dialect.compile_insert(entity, &command).map_err(Into::into)
500}
501
502fn compile_initial_graph_update(
503 dialect: &impl SqlDialect,
504 entity: &EntityDescriptor,
505 graph: &crate::GraphNode,
506) -> Result<Option<CompiledQuery>, MutationExecutorError> {
507 let Some(id) = graph.values.get("id") else {
508 return Ok(None);
509 };
510 let mut command = UpdateCommand::new(&graph.entity, id.clone());
511 for (field, value) in &graph.values {
512 if field != "id" {
513 command = command.value(field.clone(), value.clone());
514 }
515 }
516 match dialect.compile_update(entity, &command) {
517 Ok(query) => Ok(Some(query)),
518 Err(SqlCompileError::EmptyMutation(_)) => Ok(None),
519 Err(err) => Err(err.into()),
520 }
521}
522
523pub trait PostgresSchemaExt {
524 fn ensure_postgres_schema(
525 &self,
526 ) -> Pin<Box<dyn Future<Output = Result<(), MutationExecutorError>> + '_>>;
527}
528
529pub async fn ensure_postgres_schema_for(
530 context: &UserContext,
531) -> Result<(), MutationExecutorError> {
532 let dialect = context.get_resource::<PostgresDialect>().ok_or_else(|| {
533 MutationExecutorError::Bind("missing typed resource: PostgresDialect".to_owned())
534 })?;
535 let executor = context
536 .get_resource::<PgMutationExecutor>()
537 .ok_or_else(|| {
538 MutationExecutorError::Bind("missing typed resource: PgMutationExecutor".to_owned())
539 })?;
540
541 let entities = context.all_entities();
542
543 executor.ensure_schema(dialect, &entities).await?;
544 ensure_initial_graphs_postgres(executor, dialect, context).await
545}
546
547#[cfg(test)]
548mod streaming_tests {
549 use super::*;
550 use futures_util::StreamExt;
551 use teaql_sql::{SqlTransport, StreamingSqlTransport};
552
553 #[tokio::test]
554 async fn streams_from_real_postgres_when_configured() {
555 let Ok(url) = std::env::var("TEAQL_TEST_POSTGRES_URL") else {
556 return;
557 };
558 let mut config = deadpool_postgres::Config::new();
559 config.url = Some(url);
560 let pool = config
561 .create_pool(
562 Some(deadpool_postgres::Runtime::Tokio1),
563 tokio_postgres::NoTls,
564 )
565 .unwrap();
566 let executor = PgMutationExecutor::new(pool);
567 let query = CompiledQuery {
568 sql: "SELECT id FROM (VALUES (1), (2), (3), (4), (5)) AS fixture(id) ORDER BY id"
569 .to_owned(),
570 params: vec![],
571 comment: None,
572 };
573 let mut stream = executor.stream_sql(query, 2);
574 let mut sizes = Vec::new();
575 while let Some(chunk) = stream.next().await {
576 sizes.push(chunk.unwrap().rows.len());
577 }
578 assert_eq!(sizes, vec![2, 2, 1]);
579 }
580
581 #[tokio::test]
582 async fn boolean_roundtrips_real_postgres_when_configured() {
583 let Ok(url) = std::env::var("TEAQL_TEST_POSTGRES_URL") else {
584 return;
585 };
586 let mut config = deadpool_postgres::Config::new();
587 config.url = Some(url);
588 let pool = config
589 .create_pool(
590 Some(deadpool_postgres::Runtime::Tokio1),
591 tokio_postgres::NoTls,
592 )
593 .unwrap();
594 let executor = PgMutationExecutor::new(pool);
595 executor
596 .execute_sql(&CompiledQuery {
597 sql: "DROP TABLE IF EXISTS teaql_boolean_runtime_fixture".to_owned(),
598 params: vec![],
599 comment: None,
600 })
601 .await
602 .unwrap();
603 executor
604 .execute_sql(&CompiledQuery {
605 sql: "CREATE TABLE teaql_boolean_runtime_fixture(id BIGINT, required_flag BOOLEAN NOT NULL, optional_flag BOOLEAN)".to_owned(),
606 params: vec![],
607 comment: None,
608 })
609 .await
610 .unwrap();
611 for (id, required_flag, optional_flag) in [
612 (1_i64, Value::Bool(false), Value::Bool(true)),
613 (2_i64, Value::Bool(true), Value::Bool(false)),
614 (3_i64, Value::Bool(true), Value::Null),
615 ] {
616 executor
617 .execute_sql(&CompiledQuery {
618 sql: "INSERT INTO teaql_boolean_runtime_fixture VALUES ($1, $2, $3)".to_owned(),
619 params: vec![Value::I64(id), required_flag, optional_flag],
620 comment: None,
621 })
622 .await
623 .unwrap();
624 }
625 let rows = executor
626 .fetch_all_sql(&CompiledQuery {
627 sql: "SELECT required_flag, optional_flag FROM teaql_boolean_runtime_fixture ORDER BY id".to_owned(),
628 params: vec![],
629 comment: None,
630 })
631 .await
632 .unwrap();
633 assert_eq!(rows[0].get("required_flag"), Some(&Value::Bool(false)));
634 assert_eq!(rows[0].get("optional_flag"), Some(&Value::Bool(true)));
635 assert_eq!(rows[1].get("required_flag"), Some(&Value::Bool(true)));
636 assert_eq!(rows[1].get("optional_flag"), Some(&Value::Bool(false)));
637 assert_eq!(rows[2].get("optional_flag"), Some(&Value::Null));
638 executor
639 .execute_sql(&CompiledQuery {
640 sql: "DROP TABLE teaql_boolean_runtime_fixture".to_owned(),
641 params: vec![],
642 comment: None,
643 })
644 .await
645 .unwrap();
646 }
647
648 #[tokio::test]
649 async fn temporal_debug_sql_matches_real_postgres_when_configured() {
650 let Ok(url) = std::env::var("TEAQL_TEST_POSTGRES_URL") else {
651 return;
652 };
653 let mut config = deadpool_postgres::Config::new();
654 config.url = Some(url);
655 let pool = config
656 .create_pool(
657 Some(deadpool_postgres::Runtime::Tokio1),
658 tokio_postgres::NoTls,
659 )
660 .unwrap();
661 let executor = PgMutationExecutor::new(pool);
662 executor
663 .execute_sql(&CompiledQuery {
664 sql: "DROP TABLE IF EXISTS teaql_temporal_runtime_fixture".to_owned(),
665 params: vec![],
666 comment: None,
667 })
668 .await
669 .unwrap();
670 executor.execute_sql(&CompiledQuery { sql: "CREATE TABLE teaql_temporal_runtime_fixture(id BIGINT, d DATE, t TIMESTAMPTZ(3), t_local TIMESTAMP(3))".to_owned(), params: vec![], comment: None }).await.unwrap();
671 let prepared = CompiledQuery {
672 sql: "INSERT INTO teaql_temporal_runtime_fixture VALUES ($1, $2, $3, TIMESTAMP '1960-01-02 03:04:05.678')".to_owned(),
673 params: vec![
674 Value::I64(1),
675 Value::Date("2024-02-29".parse().unwrap()),
676 Value::Timestamp(teaql_core::time::Timestamp(-315_521_754_322)),
677 ],
678 comment: Some("teaql source=temporal.verify $1".to_owned()),
679 };
680 executor.execute_sql(&prepared).await.unwrap();
681 executor
682 .execute_sql(&CompiledQuery {
683 sql: prepared
684 .debug_sql(DatabaseKind::PostgreSql)
685 .replace("VALUES (1,", "VALUES (2,"),
686 params: vec![],
687 comment: None,
688 })
689 .await
690 .unwrap();
691 let rows = executor
692 .fetch_all_sql(&CompiledQuery {
693 sql: "SELECT d, t, t_local FROM teaql_temporal_runtime_fixture ORDER BY id".to_owned(),
694 params: vec![],
695 comment: None,
696 })
697 .await
698 .unwrap();
699 assert_eq!(rows[0], rows[1]);
700 executor
701 .execute_sql(&CompiledQuery {
702 sql: "DROP TABLE teaql_temporal_runtime_fixture".to_owned(),
703 params: vec![],
704 comment: None,
705 })
706 .await
707 .unwrap();
708 }
709}
710
711impl PostgresSchemaExt for UserContext {
712 fn ensure_postgres_schema(
713 &self,
714 ) -> Pin<Box<dyn Future<Output = Result<(), MutationExecutorError>> + '_>> {
715 Box::pin(ensure_postgres_schema_for(self))
716 }
717}
718
719#[derive(Debug, Default, Clone, Copy)]
720pub struct PostgresSchemaProvider;
721
722impl SchemaProvider for PostgresSchemaProvider {
723 fn ensure_schema<'a>(
724 &'a self,
725 context: &'a UserContext,
726 ) -> Pin<Box<dyn Future<Output = Result<(), RuntimeError>> + Send + 'a>> {
727 Box::pin(async move {
728 ensure_postgres_schema_for(context)
729 .await
730 .map_err(|err| RuntimeError::Schema(err.to_string()))
731 })
732 }
733}
734
735pub trait PostgresProviderExt {
736 fn use_postgres_provider(&mut self, executor: PgMutationExecutor) -> &mut Self;
737}
738
739impl PostgresProviderExt for UserContext {
740 fn use_postgres_provider(&mut self, executor: PgMutationExecutor) -> &mut Self {
741 self.insert_resource(PostgresDialect);
742 self.insert_resource(executor);
743 self.set_schema_provider(PostgresSchemaProvider);
744 self
745 }
746}
747
748#[derive(Clone)]
749pub struct PgIdSpaceGenerator {
750 pool: Pool,
751 table_name: String,
752}
753
754impl PgIdSpaceGenerator {
755 pub fn new(pool: Pool) -> Self {
756 Self {
757 pool,
758 table_name: DEFAULT_ID_SPACE_TABLE.to_owned(),
759 }
760 }
761
762 pub fn from_executor(executor: PgMutationExecutor) -> Self {
763 Self::new(executor.pool())
764 }
765
766 pub fn with_table_name(mut self, table_name: impl Into<String>) -> Self {
767 self.table_name = table_name.into();
768 self
769 }
770
771 pub async fn ensure_table(&self) -> Result<(), MutationExecutorError> {
772 PgMutationExecutor::new(self.pool.clone())
773 .ensure_id_space_table(&self.table_name)
774 .await
775 }
776
777 pub async fn next_id(&self, entity: &str) -> Result<u64, MutationExecutorError> {
778 self.ensure_table().await?;
779 let table = quote_ident(&self.table_name);
780 let client = self
781 .pool
782 .get()
783 .await
784 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
785 let select_sql = format!("SELECT current_level FROM {table} WHERE type_name = $1");
786 let insert_sql = format!("INSERT INTO {table}(type_name, current_level) VALUES ($1, 1)");
787 let update_sql = format!("UPDATE {table} SET current_level = $1 WHERE type_name = $2 AND current_level = $3");
788 for _ in 1..=100 {
789 let current = client.query_opt(&select_sql, &[&entity]).await?
790 .map(|row| row.try_get::<_, i64>(0)).transpose()?;
791 if let Some(current) = current {
792 let next = current.checked_add(1).ok_or_else(|| MutationExecutorError::Bind(
793 format!("ID space overflow for {entity}")))?;
794 if client.execute(&update_sql, &[&next, &entity, ¤t]).await? == 1 {
795 return u64::try_from(next).map_err(|_| MutationExecutorError::Bind(
796 format!("generated id {next} cannot be represented as u64")));
797 }
798 } else {
799 match client.execute(&insert_sql, &[&entity]).await {
800 Ok(1) => return Ok(1),
801 Ok(changed) => return Err(MutationExecutorError::Bind(
802 format!("ID space insert for {entity} changed {changed} rows"))),
803 Err(error) => {
804 if client.query_opt(&select_sql, &[&entity]).await?.is_none() {
805 return Err(error.into());
806 }
807 }
808 }
809 }
810 }
811 Err(MutationExecutorError::Bind(format!(
812 "Unable to allocate ID for {entity} after 100 optimistic-lock attempts")))
813 }
814
815 pub async fn ensure_floor(&self, entity: &str, floor: u64) -> Result<(), MutationExecutorError> {
816 self.ensure_table().await?;
817 let floor = i64::try_from(floor).map_err(|_| MutationExecutorError::Bind(
818 format!("ID space floor {floor} for {entity} exceeds BIGINT")))?;
819 let table = quote_ident(&self.table_name);
820 let client = self.pool.get().await.map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
821 let select = format!("SELECT current_level FROM {table} WHERE type_name = $1");
822 let insert = format!("INSERT INTO {table}(type_name, current_level) VALUES ($1, $2)");
823 let update = format!("UPDATE {table} SET current_level = $1 WHERE type_name = $2 AND current_level = $3");
824 for _ in 1..=100 {
825 let current = client.query_opt(&select, &[&entity]).await?
826 .map(|row| row.try_get::<_, i64>(0)).transpose()?;
827 match current {
828 Some(current) if current >= floor => return Ok(()),
829 Some(current) => if client.execute(&update, &[&floor, &entity, ¤t]).await? == 1 { return Ok(()); },
830 None => match client.execute(&insert, &[&entity, &floor]).await {
831 Ok(1) => return Ok(()),
832 Ok(_) => {}
833 Err(error) => if client.query_opt(&select, &[&entity]).await?.is_none() { return Err(error.into()); },
834 },
835 }
836 }
837 Err(MutationExecutorError::Bind(format!(
838 "Unable to synchronize ID space floor for {entity} after 100 optimistic-lock attempts")))
839 }
840}
841
842impl InternalIdGenerator for PgIdSpaceGenerator {
843 fn generate_id(&self, entity: &str) -> Result<u64, RuntimeError> {
844 let generator = self.clone();
845 let entity = entity.to_owned();
846 block_on_id_generation(async move { generator.next_id(&entity).await })
847 }
848}
849
850fn block_on_id_generation<F>(future: F) -> Result<u64, RuntimeError>
851where
852 F: Future<Output = Result<u64, MutationExecutorError>> + Send + 'static,
853{
854 let result = match tokio::runtime::Handle::try_current() {
855 Ok(handle) => tokio::task::block_in_place(|| handle.block_on(future)),
856 Err(_) => tokio::runtime::Builder::new_current_thread()
857 .enable_all()
858 .build()
859 .map_err(|err| RuntimeError::IdGeneration(err.to_string()))?
860 .block_on(future),
861 };
862 result.map_err(|err| RuntimeError::IdGeneration(err.to_string()))
863}
864
865fn quote_ident(ident: &str) -> String {
866 quote_identifier_if_needed(ident, '"')
867}
868
869fn strip_identifier_quotes(ident: &str) -> &str {
873 let bytes = ident.as_bytes();
874 if bytes.len() >= 2 {
875 let (first, last) = (bytes[0], bytes[bytes.len() - 1]);
876 if (first == b'"' && last == b'"')
877 || (first == b'`' && last == b'`')
878 || (first == b'[' && last == b']')
879 {
880 return &ident[1..ident.len() - 1];
881 }
882 }
883 ident
884}
885
886fn try_parse_datetime_from_str(s: &str) -> Option<chrono::DateTime<chrono::Utc>> {
887 if let Ok(dt) = chrono::DateTime::parse_from_rfc3339(s) {
888 return Some(dt.with_timezone(&chrono::Utc));
889 }
890 if let Ok(ndt) = chrono::NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
891 return Some(chrono::DateTime::from_naive_utc_and_offset(
892 ndt,
893 chrono::Utc,
894 ));
895 }
896 if let Ok(nd) = chrono::NaiveDate::parse_from_str(s, "%Y-%m-%d") {
897 let ndt = nd.and_hms_opt(0, 0, 0)?;
898 return Some(chrono::DateTime::from_naive_utc_and_offset(
899 ndt,
900 chrono::Utc,
901 ));
902 }
903 None
904}
905
906#[derive(Debug, Clone, Copy)]
907struct PgNull;
908
909impl tokio_postgres::types::ToSql for PgNull {
910 fn to_sql(
911 &self,
912 ty: &tokio_postgres::types::Type,
913 out: &mut bytes::BytesMut,
914 ) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>> {
915 Ok(tokio_postgres::types::IsNull::Yes)
916 }
917
918 fn accepts(ty: &tokio_postgres::types::Type) -> bool {
919 true
920 }
921
922 fn to_sql_checked(
923 &self,
924 ty: &tokio_postgres::types::Type,
925 out: &mut bytes::BytesMut,
926 ) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>> {
927 Ok(tokio_postgres::types::IsNull::Yes)
928 }
929}
930
931#[derive(Debug, Clone, Copy)]
932struct PgTimestamp(DateTime<Utc>);
933
934impl tokio_postgres::types::ToSql for PgTimestamp {
935 fn to_sql(
936 &self,
937 ty: &tokio_postgres::types::Type,
938 out: &mut bytes::BytesMut,
939 ) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>> {
940 if *ty == tokio_postgres::types::Type::TIMESTAMP {
941 self.0.naive_utc().to_sql(ty, out)
942 } else {
943 self.0.to_sql(ty, out)
944 }
945 }
946
947 fn accepts(ty: &tokio_postgres::types::Type) -> bool {
948 *ty == tokio_postgres::types::Type::TIMESTAMP
949 || *ty == tokio_postgres::types::Type::TIMESTAMPTZ
950 }
951
952 tokio_postgres::types::to_sql_checked!();
953}
954
955struct PgArgs {
956 values: Vec<Box<dyn tokio_postgres::types::ToSql + Sync + Send>>,
957}
958impl PgArgs {
959 fn add<T: tokio_postgres::types::ToSql + Sync + Send + 'static>(&mut self, v: T) {
960 self.values.push(Box::new(v));
961 }
962 fn as_refs(&self) -> Vec<&(dyn tokio_postgres::types::ToSql + Sync)> {
963 self.values.iter().map(|b| b.as_ref() as _).collect()
964 }
965}
966
967fn bind_pg(args: &mut PgArgs, value: &Value) -> Result<(), MutationExecutorError> {
968 match value {
969 Value::Null => {
970 args.add(PgNull);
971 }
972 Value::Bool(v) => args.add(*v),
973 Value::I64(v) => args.add(*v),
974 Value::U64(v) => {
975 let v = i64::try_from(*v).map_err(|_| {
976 MutationExecutorError::Bind(format!("u64 value {v} exceeds i64 range"))
977 })?;
978 args.add(v);
979 }
980 Value::F64(v) => args.add(*v),
981 Value::Decimal(v) => args.add(*v),
982 Value::Text(v) => match try_parse_datetime_from_str(v) {
983 Some(dt) => args.add(dt),
984 None => args.add(v.clone()),
985 },
986 Value::Json(v) => {
987 let j_val: serde_json::Value =
988 serde_json::to_value(v).map_err(|e| MutationExecutorError::Bind(e.to_string()))?;
989 args.add(j_val);
990 }
991 Value::Date(v) => args.add(*v),
992 Value::Timestamp(v) => args.add(PgTimestamp(v.to_datetime())),
993 Value::Object(_) => return Err(MutationExecutorError::UnsupportedValue("object")),
994 Value::List(values) => bind_pg_list(args, values)?,
995 Value::TypedNull(dt) => match dt {
996 DataType::Bool => args.add(Option::<bool>::None),
997 DataType::I64 | DataType::U64 => args.add(Option::<i64>::None),
998 DataType::F64 => args.add(Option::<f64>::None),
999 DataType::Decimal => args.add(Option::<Decimal>::None),
1000 DataType::Text | DataType::LargeText => args.add(Option::<String>::None),
1001 DataType::Json => args.add(Option::<serde_json::Value>::None),
1002 DataType::Date => args.add(Option::<NaiveDate>::None),
1003 DataType::Timestamp => args.add(PgNull),
1004 },
1005 }
1006 Ok(())
1007}
1008
1009fn bind_pg_list(args: &mut PgArgs, values: &[Value]) -> Result<(), MutationExecutorError> {
1010 let Some(first) = values.first() else {
1011 return Err(MutationExecutorError::UnsupportedValue("empty list"));
1012 };
1013 match first {
1014 Value::Bool(_) => {
1015 let values = values
1016 .iter()
1017 .map(|value| match value {
1018 Value::Bool(value) => Ok(*value),
1019 _ => Err(MutationExecutorError::UnsupportedValue("mixed bool list")),
1020 })
1021 .collect::<Result<Vec<_>, _>>()?;
1022 args.add(values);
1023 }
1024 Value::I64(_) => {
1025 let values = values
1026 .iter()
1027 .map(|value| match value {
1028 Value::I64(value) => Ok(*value),
1029 _ => Err(MutationExecutorError::UnsupportedValue("mixed i64 list")),
1030 })
1031 .collect::<Result<Vec<_>, _>>()?;
1032 args.add(values);
1033 }
1034 Value::U64(_) => {
1035 let values = values
1036 .iter()
1037 .map(|value| match value {
1038 Value::U64(value) => i64::try_from(*value).map_err(|_| {
1039 MutationExecutorError::Bind(format!("u64 value {value} exceeds i64 range"))
1040 }),
1041 _ => Err(MutationExecutorError::UnsupportedValue("mixed u64 list")),
1042 })
1043 .collect::<Result<Vec<_>, _>>()?;
1044 args.add(values);
1045 }
1046 Value::F64(_) => {
1047 let values = values
1048 .iter()
1049 .map(|value| match value {
1050 Value::F64(value) => Ok(*value),
1051 _ => Err(MutationExecutorError::UnsupportedValue("mixed f64 list")),
1052 })
1053 .collect::<Result<Vec<_>, _>>()?;
1054 args.add(values);
1055 }
1056 Value::Decimal(_) => {
1057 let values = values
1058 .iter()
1059 .map(|value| match value {
1060 Value::Decimal(value) => Ok(*value),
1061 _ => Err(MutationExecutorError::UnsupportedValue(
1062 "mixed decimal list",
1063 )),
1064 })
1065 .collect::<Result<Vec<_>, _>>()?;
1066 args.add(values);
1067 }
1068 Value::Text(_) => {
1069 let values = values
1070 .iter()
1071 .map(|value| match value {
1072 Value::Text(value) => Ok(value.clone()),
1073 _ => Err(MutationExecutorError::UnsupportedValue("mixed text list")),
1074 })
1075 .collect::<Result<Vec<_>, _>>()?;
1076 args.add(values);
1077 }
1078 Value::Date(_) => {
1079 let values = values
1080 .iter()
1081 .map(|value| match value {
1082 Value::Date(value) => Ok(*value),
1083 _ => Err(MutationExecutorError::UnsupportedValue("mixed date list")),
1084 })
1085 .collect::<Result<Vec<_>, _>>()?;
1086 args.add(values);
1087 }
1088 Value::Timestamp(_) => {
1089 let values = values
1090 .iter()
1091 .map(|value| match value {
1092 Value::Timestamp(value) => Ok(value.to_datetime()),
1093 _ => Err(MutationExecutorError::UnsupportedValue(
1094 "mixed timestamp list",
1095 )),
1096 })
1097 .collect::<Result<Vec<_>, _>>()?;
1098 args.add(values);
1099 }
1100 Value::Null => return Err(MutationExecutorError::UnsupportedValue("null list")),
1101 Value::Json(_) => return Err(MutationExecutorError::UnsupportedValue("json list")),
1102 Value::Object(_) => return Err(MutationExecutorError::UnsupportedValue("object list")),
1103 Value::List(_) => return Err(MutationExecutorError::UnsupportedValue("nested list")),
1104 Value::TypedNull(_) => return Err(MutationExecutorError::UnsupportedValue("null list")),
1105 }
1106 Ok(())
1107}
1108
1109fn decode_pg_row(row: &tokio_postgres::Row) -> Result<Record, MutationExecutorError> {
1110 let mut record = BTreeMap::new();
1111 for (index, column) in row.columns().iter().enumerate() {
1112 let name = column.name().to_owned();
1113 let type_name = column.type_().name().to_ascii_uppercase();
1114
1115 let value = match type_name.as_str() {
1116 "BOOL" | "BOOLEAN" => {
1117 let v: Option<bool> = row.try_get(index)?;
1118 match v {
1119 Some(v) => Value::Bool(v),
1120 None => Value::Null,
1121 }
1122 }
1123 "INT2" => {
1124 let v: Option<i16> = row.try_get(index)?;
1125 match v {
1126 Some(v) => Value::I64(v as i64),
1127 None => Value::Null,
1128 }
1129 }
1130 "INT4" => {
1131 let v: Option<i32> = row.try_get(index)?;
1132 match v {
1133 Some(v) => Value::I64(v as i64),
1134 None => Value::Null,
1135 }
1136 }
1137 "INT8" => {
1138 let v: Option<i64> = row.try_get(index)?;
1139 match v {
1140 Some(v) => Value::I64(v),
1141 None => Value::Null,
1142 }
1143 }
1144 "FLOAT4" => {
1145 let v: Option<f32> = row.try_get(index)?;
1146 match v {
1147 Some(v) => Value::F64(v as f64),
1148 None => Value::Null,
1149 }
1150 }
1151 "FLOAT8" => {
1152 let v: Option<f64> = row.try_get(index)?;
1153 match v {
1154 Some(v) => Value::F64(v),
1155 None => Value::Null,
1156 }
1157 }
1158 "NUMERIC" => {
1159 let v: Option<Decimal> = row.try_get(index)?;
1160 match v {
1161 Some(v) => Value::Decimal(v),
1162 None => Value::Null,
1163 }
1164 }
1165 "JSON" | "JSONB" => {
1166 let v: Option<serde_json::Value> = row.try_get(index)?;
1167 match v {
1168 Some(j) => Value::Json(j.into()),
1169 None => Value::Null,
1170 }
1171 }
1172 "DATE" => {
1173 let v: Option<NaiveDate> = row.try_get(index)?;
1174 match v {
1175 Some(v) => Value::Date(v),
1176 None => Value::Null,
1177 }
1178 }
1179 "TIMESTAMP" => {
1180 let v: Option<NaiveDateTime> = row.try_get(index)?;
1181 match v {
1182 Some(v) => Value::Timestamp(teaql_core::time::Timestamp(
1183 v.and_utc().timestamp_millis(),
1184 )),
1185 None => Value::Null,
1186 }
1187 }
1188 "TIMESTAMPTZ" => {
1189 let v: Option<DateTime<Utc>> = row.try_get(index)?;
1190 match v {
1191 Some(v) => Value::Timestamp(teaql_core::time::Timestamp(v.timestamp_millis())),
1192 None => Value::Null,
1193 }
1194 }
1195 "TEXT" | "VARCHAR" | "BPCHAR" | "NAME" | "UUID" => {
1196 let v: Option<String> = row.try_get(index)?;
1197 match v {
1198 Some(v) => Value::Text(v),
1199 None => Value::Null,
1200 }
1201 }
1202 other => {
1203 return Err(MutationExecutorError::UnsupportedColumnType(
1204 other.to_owned(),
1205 ));
1206 }
1207 };
1208 record.insert(name, value);
1209 }
1210 Ok(record)
1211}
1212
1213#[cfg(test)]
1214mod tests {
1215 use super::*;
1216 use teaql_core::{DeleteCommand, RecoverCommand};
1217
1218 fn entity() -> EntityDescriptor {
1219 EntityDescriptor::new("Order")
1220 .table_name("orders")
1221 .property(
1222 PropertyDescriptor::new("id", DataType::U64)
1223 .column_name("id")
1224 .id()
1225 .not_null(),
1226 )
1227 .property(
1228 PropertyDescriptor::new("version", DataType::I64)
1229 .column_name("version")
1230 .version()
1231 .not_null(),
1232 )
1233 .property(PropertyDescriptor::new("name", DataType::Text).column_name("name"))
1234 }
1235
1236 #[test]
1237 fn postgres_dialect_compiles_mutations_with_numbered_placeholders() {
1238 let insert = PostgresDialect
1239 .compile_insert(
1240 &entity(),
1241 &InsertCommand::new("Order")
1242 .value("id", 1_u64)
1243 .value("name", "A"),
1244 )
1245 .unwrap();
1246 assert_eq!(insert.sql, "INSERT INTO orders (id, name) VALUES ($1, $2)");
1247
1248 let update = PostgresDialect
1249 .compile_update(
1250 &entity(),
1251 &UpdateCommand::new("Order", 1_u64)
1252 .expected_version(3)
1253 .value("name", "B"),
1254 )
1255 .unwrap();
1256 assert_eq!(
1257 update.sql,
1258 "UPDATE orders SET name = $1, version = $2 WHERE id = $3 AND version = $4"
1259 );
1260
1261 let delete = PostgresDialect
1262 .compile_delete(
1263 &entity(),
1264 &DeleteCommand::new("Order", 1_u64).expected_version(3),
1265 )
1266 .unwrap();
1267 let recover = PostgresDialect
1268 .compile_recover(&entity(), &RecoverCommand::new("Order", 1_u64, -4))
1269 .unwrap();
1270 assert_eq!(
1271 delete.sql,
1272 "UPDATE orders SET version = $1 WHERE id = $2 AND version = $3"
1273 );
1274 assert_eq!(
1275 recover.sql,
1276 "UPDATE orders SET version = $1 WHERE id = $2 AND version = $3"
1277 );
1278 }
1279
1280 #[test]
1281 fn postgres_dialect_compiles_schema_and_large_in_array_binds() {
1282 let create = PostgresDialect.compile_create_table(&entity()).unwrap();
1283 assert_eq!(
1284 create,
1285 "CREATE TABLE IF NOT EXISTS orders (id BIGINT PRIMARY KEY NOT NULL, version BIGINT NOT NULL, name VARCHAR(255))"
1286 );
1287 assert!(
1288 PostgresDialect
1289 .schema_setup_sqls()
1290 .iter()
1291 .any(|sql| sql.contains("CREATE OR REPLACE FUNCTION soundex"))
1292 );
1293
1294 let query = PostgresDialect
1295 .compile_select(
1296 &entity(),
1297 &SelectQuery::new("Order")
1298 .filter(Expr::in_large(
1299 "id",
1300 vec![Value::from(1_u64), Value::from(2_u64)],
1301 ))
1302 .order_asc("id"),
1303 )
1304 .unwrap();
1305 assert_eq!(
1306 query.sql,
1307 "SELECT id, version, name FROM orders WHERE (id = ANY($1)) ORDER BY id ASC"
1308 );
1309 assert_eq!(
1310 query.params,
1311 vec![Value::List(vec![Value::U64(1), Value::U64(2)])]
1312 );
1313 }
1314}