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