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