1#![allow(warnings)]
2use std::collections::BTreeMap;
3use std::future::Future;
4use std::pin::Pin;
5
6use chrono::{DateTime, NaiveDate, 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 rows = client.query_raw(&query.sql, params).await?;
252 futures_util::pin_mut!(rows);
253 let mut chunk = Vec::with_capacity(chunk_size); let mut index = 0;
254 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; } }
255 if !chunk.is_empty() { yield teaql_data_service::StreamChunk { rows:chunk, chunk_index:index, is_last:true }; }
256 })
257 }
258}
259
260impl teaql_sql::SqlTransaction for PgMutationExecutor {
261 type Error = MutationExecutorError;
262
263 async fn commit_sql(self) -> Result<(), Self::Error> {
264 Err(MutationExecutorError::Bind(
265 "Transactions not supported yet".to_string(),
266 ))
267 }
268
269 async fn rollback_sql(self) -> Result<(), Self::Error> {
270 Err(MutationExecutorError::Bind(
271 "Transactions not supported yet".to_string(),
272 ))
273 }
274}
275
276impl teaql_sql::SqlTransactionTransport for PgMutationExecutor {
277 type Tx<'a>
278 = Self
279 where
280 Self: 'a;
281
282 async fn begin_sql(&self) -> Result<Self::Tx<'_>, Self::Error> {
283 Err(MutationExecutorError::Bind(
284 "Transactions not supported yet".to_string(),
285 ))
286 }
287}
288
289impl PgMutationExecutor {
290 pub fn new(pool: Pool) -> Self {
291 Self { pool }
292 }
293
294 pub fn pool(&self) -> Pool {
295 self.pool.clone()
296 }
297
298 pub async fn ensure_schema(
299 &self,
300 dialect: &PostgresDialect,
301 entities: &[&EntityDescriptor],
302 ) -> Result<(), MutationExecutorError> {
303 let client = self
304 .pool
305 .get()
306 .await
307 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
308 for sql in dialect.schema_setup_sqls() {
309 client.execute(*sql, &[]).await?;
310 }
311 self.ensure_id_space_table(DEFAULT_ID_SPACE_TABLE).await?;
312
313 for entity in entities {
314 if !self.table_exists(&entity.table_name).await? {
315 let sql = dialect.compile_create_table(entity)?;
316 client.execute(&sql, &[]).await?;
317 continue;
318 }
319
320 let existing_columns = self.table_columns(&entity.table_name).await?;
321 for property in &entity.properties {
322 let bare_column = strip_identifier_quotes(&property.column_name).to_lowercase();
323 if existing_columns.contains(&bare_column) {
324 continue;
325 }
326 let sql = dialect.compile_add_column(entity, property)?;
327 client.execute(&sql, &[]).await?;
328 }
329
330 for sql in dialect.schema_indexes_sqls(entity)? {
331 client.execute(&sql, &[]).await?;
332 }
333 }
334 Ok(())
335 }
336
337 pub async fn ensure_id_space_table(
338 &self,
339 table_name: &str,
340 ) -> Result<(), MutationExecutorError> {
341 let sql = format!(
342 "CREATE TABLE IF NOT EXISTS {} (type_name VARCHAR(100) PRIMARY KEY, current_level BIGINT NOT NULL)",
343 quote_ident(table_name)
344 );
345 let client = self
346 .pool
347 .get()
348 .await
349 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
350 client.execute(&sql, &[]).await?;
351 Ok(())
352 }
353
354 pub async fn execute(&self, query: &CompiledQuery) -> Result<u64, MutationExecutorError> {
355 let mut args = PgArgs { values: Vec::new() };
356 for value in &query.params {
357 bind_pg(&mut args, value)?;
358 }
359 let client = self
360 .pool
361 .get()
362 .await
363 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
364 let result = client.execute(&query.sql, &args.as_refs()).await?;
365 Ok(result)
366 }
367
368 pub async fn fetch_all(
369 &self,
370 query: &CompiledQuery,
371 ) -> Result<Vec<Record>, MutationExecutorError> {
372 let mut args = PgArgs { values: Vec::new() };
373 for value in &query.params {
374 bind_pg(&mut args, value)?;
375 }
376 let client = self
377 .pool
378 .get()
379 .await
380 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
381 let rows = client.query(&query.sql, &args.as_refs()).await?;
382 rows.iter().map(decode_pg_row).collect()
383 }
384
385 async fn table_exists(&self, table_name: &str) -> Result<bool, MutationExecutorError> {
386 let client = self
387 .pool
388 .get()
389 .await
390 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
391 let row = client
392 .query_one(
393 "SELECT COUNT(1)
394 FROM information_schema.tables
395 WHERE table_schema = current_schema()
396 AND table_name = $1",
397 &[&table_name],
398 )
399 .await?;
400 let exists: i64 = row.try_get(0)?;
401 Ok(exists > 0)
402 }
403
404 async fn table_columns(
405 &self,
406 table_name: &str,
407 ) -> Result<std::collections::BTreeSet<String>, MutationExecutorError> {
408 let client = self
409 .pool
410 .get()
411 .await
412 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
413 let rows = client
414 .query(
415 "SELECT column_name
416 FROM information_schema.columns
417 WHERE table_schema = current_schema()
418 AND table_name = $1",
419 &[&table_name],
420 )
421 .await?;
422 let mut columns = std::collections::BTreeSet::new();
423 for row in rows {
424 let name: String = row.try_get("column_name")?;
425 columns.insert(name.to_lowercase());
426 }
427 Ok(columns)
428 }
429}
430
431async fn ensure_initial_graphs_postgres(
432 executor: &PgMutationExecutor,
433 dialect: &PostgresDialect,
434 context: &UserContext,
435) -> Result<(), MutationExecutorError> {
436 for graph in context.initial_graphs() {
437 let entity = context.entity(&graph.entity).ok_or_else(|| {
438 MutationExecutorError::Bind(format!("missing entity: {}", graph.entity))
439 })?;
440 if initial_graph_exists_postgres(executor, dialect, entity, graph).await? {
441 if let Some(query) = compile_initial_graph_update(dialect, entity, graph)? {
442 executor.execute(&query).await?;
443 }
444 continue;
445 }
446 let query = compile_initial_graph_insert(dialect, entity, graph)?;
447 executor.execute(&query).await?;
448 }
449 for graph in context.root_graphs() {
450 let entity = context.entity(&graph.entity).ok_or_else(|| {
451 MutationExecutorError::Bind(format!("missing entity: {}", graph.entity))
452 })?;
453 if initial_graph_exists_postgres(executor, dialect, entity, graph).await? {
454 continue;
455 }
456 let query = compile_initial_graph_insert(dialect, entity, graph)?;
457 executor.execute(&query).await?;
458 }
459 Ok(())
460}
461
462async fn initial_graph_exists_postgres(
463 executor: &PgMutationExecutor,
464 dialect: &PostgresDialect,
465 entity: &EntityDescriptor,
466 graph: &GraphNode,
467) -> Result<bool, MutationExecutorError> {
468 let Some(id) = graph.values.get("id") else {
469 return Ok(false);
470 };
471 let query = dialect.compile_select(
472 entity,
473 &SelectQuery::new(&graph.entity)
474 .project("id")
475 .filter(Expr::eq("id", id.clone()))
476 .limit(1),
477 )?;
478 Ok(!executor.fetch_all(&query).await?.is_empty())
479}
480
481fn compile_initial_graph_insert(
482 dialect: &impl SqlDialect,
483 entity: &EntityDescriptor,
484 graph: &GraphNode,
485) -> Result<CompiledQuery, MutationExecutorError> {
486 let mut command = InsertCommand::new(&graph.entity);
487 for (field, value) in &graph.values {
488 command = command.value(field.clone(), value.clone());
489 }
490 dialect.compile_insert(entity, &command).map_err(Into::into)
491}
492
493fn compile_initial_graph_update(
494 dialect: &impl SqlDialect,
495 entity: &EntityDescriptor,
496 graph: &crate::GraphNode,
497) -> Result<Option<CompiledQuery>, MutationExecutorError> {
498 let Some(id) = graph.values.get("id") else {
499 return Ok(None);
500 };
501 let mut command = UpdateCommand::new(&graph.entity, id.clone());
502 for (field, value) in &graph.values {
503 if field != "id" {
504 command = command.value(field.clone(), value.clone());
505 }
506 }
507 match dialect.compile_update(entity, &command) {
508 Ok(query) => Ok(Some(query)),
509 Err(SqlCompileError::EmptyMutation(_)) => Ok(None),
510 Err(err) => Err(err.into()),
511 }
512}
513
514pub trait PostgresSchemaExt {
515 fn ensure_postgres_schema(
516 &self,
517 ) -> Pin<Box<dyn Future<Output = Result<(), MutationExecutorError>> + '_>>;
518}
519
520pub async fn ensure_postgres_schema_for(
521 context: &UserContext,
522) -> Result<(), MutationExecutorError> {
523 let dialect = context.get_resource::<PostgresDialect>().ok_or_else(|| {
524 MutationExecutorError::Bind("missing typed resource: PostgresDialect".to_owned())
525 })?;
526 let executor = context
527 .get_resource::<PgMutationExecutor>()
528 .ok_or_else(|| {
529 MutationExecutorError::Bind("missing typed resource: PgMutationExecutor".to_owned())
530 })?;
531
532 let entities = context.all_entities();
533
534 executor.ensure_schema(dialect, &entities).await?;
535 ensure_initial_graphs_postgres(executor, dialect, context).await
536}
537
538#[cfg(test)]
539mod streaming_tests {
540 use super::*;
541 use futures_util::StreamExt;
542 use teaql_sql::{SqlTransport, StreamingSqlTransport};
543
544 #[tokio::test]
545 async fn streams_from_real_postgres_when_configured() {
546 let Ok(url) = std::env::var("TEAQL_TEST_POSTGRES_URL") else {
547 return;
548 };
549 let mut config = deadpool_postgres::Config::new();
550 config.url = Some(url);
551 let pool = config
552 .create_pool(
553 Some(deadpool_postgres::Runtime::Tokio1),
554 tokio_postgres::NoTls,
555 )
556 .unwrap();
557 let executor = PgMutationExecutor::new(pool);
558 let query = CompiledQuery {
559 sql: "SELECT id FROM (VALUES (1), (2), (3), (4), (5)) AS fixture(id) ORDER BY id"
560 .to_owned(),
561 params: vec![],
562 comment: None,
563 };
564 let mut stream = executor.stream_sql(query, 2);
565 let mut sizes = Vec::new();
566 while let Some(chunk) = stream.next().await {
567 sizes.push(chunk.unwrap().rows.len());
568 }
569 assert_eq!(sizes, vec![2, 2, 1]);
570 }
571
572 #[tokio::test]
573 async fn temporal_debug_sql_matches_real_postgres_when_configured() {
574 let Ok(url) = std::env::var("TEAQL_TEST_POSTGRES_URL") else {
575 return;
576 };
577 let mut config = deadpool_postgres::Config::new();
578 config.url = Some(url);
579 let pool = config
580 .create_pool(
581 Some(deadpool_postgres::Runtime::Tokio1),
582 tokio_postgres::NoTls,
583 )
584 .unwrap();
585 let executor = PgMutationExecutor::new(pool);
586 executor
587 .execute_sql(&CompiledQuery {
588 sql: "DROP TABLE IF EXISTS teaql_temporal_runtime_fixture".to_owned(),
589 params: vec![],
590 comment: None,
591 })
592 .await
593 .unwrap();
594 executor.execute_sql(&CompiledQuery { sql: "CREATE TABLE teaql_temporal_runtime_fixture(id BIGINT, d DATE, t TIMESTAMPTZ(3))".to_owned(), params: vec![], comment: None }).await.unwrap();
595 let prepared = CompiledQuery {
596 sql: "INSERT INTO teaql_temporal_runtime_fixture VALUES ($1, $2, $3)".to_owned(),
597 params: vec![
598 Value::I64(1),
599 Value::Date("2024-02-29".parse().unwrap()),
600 Value::Timestamp(teaql_core::time::Timestamp(-315_521_754_322)),
601 ],
602 comment: Some("teaql source=temporal.verify $1".to_owned()),
603 };
604 executor.execute_sql(&prepared).await.unwrap();
605 executor
606 .execute_sql(&CompiledQuery {
607 sql: prepared
608 .debug_sql(DatabaseKind::PostgreSql)
609 .replace("VALUES (1,", "VALUES (2,"),
610 params: vec![],
611 comment: None,
612 })
613 .await
614 .unwrap();
615 let rows = executor
616 .fetch_all_sql(&CompiledQuery {
617 sql: "SELECT d, t FROM teaql_temporal_runtime_fixture ORDER BY id".to_owned(),
618 params: vec![],
619 comment: None,
620 })
621 .await
622 .unwrap();
623 assert_eq!(rows[0], rows[1]);
624 executor
625 .execute_sql(&CompiledQuery {
626 sql: "DROP TABLE teaql_temporal_runtime_fixture".to_owned(),
627 params: vec![],
628 comment: None,
629 })
630 .await
631 .unwrap();
632 }
633}
634
635impl PostgresSchemaExt for UserContext {
636 fn ensure_postgres_schema(
637 &self,
638 ) -> Pin<Box<dyn Future<Output = Result<(), MutationExecutorError>> + '_>> {
639 Box::pin(ensure_postgres_schema_for(self))
640 }
641}
642
643#[derive(Debug, Default, Clone, Copy)]
644pub struct PostgresSchemaProvider;
645
646impl SchemaProvider for PostgresSchemaProvider {
647 fn ensure_schema<'a>(
648 &'a self,
649 context: &'a UserContext,
650 ) -> Pin<Box<dyn Future<Output = Result<(), RuntimeError>> + Send + 'a>> {
651 Box::pin(async move {
652 ensure_postgres_schema_for(context)
653 .await
654 .map_err(|err| RuntimeError::Schema(err.to_string()))
655 })
656 }
657}
658
659pub trait PostgresProviderExt {
660 fn use_postgres_provider(&mut self, executor: PgMutationExecutor) -> &mut Self;
661}
662
663impl PostgresProviderExt for UserContext {
664 fn use_postgres_provider(&mut self, executor: PgMutationExecutor) -> &mut Self {
665 self.insert_resource(PostgresDialect);
666 self.insert_resource(executor);
667 self.set_schema_provider(PostgresSchemaProvider);
668 self
669 }
670}
671
672#[derive(Clone)]
673pub struct PgIdSpaceGenerator {
674 pool: Pool,
675 table_name: String,
676}
677
678impl PgIdSpaceGenerator {
679 pub fn new(pool: Pool) -> Self {
680 Self {
681 pool,
682 table_name: DEFAULT_ID_SPACE_TABLE.to_owned(),
683 }
684 }
685
686 pub fn from_executor(executor: PgMutationExecutor) -> Self {
687 Self::new(executor.pool())
688 }
689
690 pub fn with_table_name(mut self, table_name: impl Into<String>) -> Self {
691 self.table_name = table_name.into();
692 self
693 }
694
695 pub async fn ensure_table(&self) -> Result<(), MutationExecutorError> {
696 PgMutationExecutor::new(self.pool.clone())
697 .ensure_id_space_table(&self.table_name)
698 .await
699 }
700
701 pub async fn next_id(&self, entity: &str) -> Result<u64, MutationExecutorError> {
702 self.ensure_table().await?;
703 let update_sql = format!(
704 "UPDATE {} SET current_level = current_level + 1 WHERE type_name = $1 RETURNING current_level",
705 quote_ident(&self.table_name)
706 );
707 let client = self
708 .pool
709 .get()
710 .await
711 .map_err(|e| MutationExecutorError::Pool(e.to_string()))?;
712 let row = client.query_opt(&update_sql, &[&entity]).await?;
713
714 let id = match row {
715 Some(r) => {
716 let level: i64 = r.try_get(0)?;
717 level
718 }
719 None => {
720 let insert_sql = format!(
721 "INSERT INTO {} (type_name, current_level) VALUES ($1, 1) RETURNING current_level",
722 quote_ident(&self.table_name)
723 );
724 let insert_res = client.query_one(&insert_sql, &[&entity]).await;
725 match insert_res {
726 Ok(r) => {
727 let level: i64 = r.try_get(0)?;
728 level
729 }
730 Err(_) => {
731 let row = client.query_one(&update_sql, &[&entity]).await?;
732 let level: i64 = row.try_get(0)?;
733 level
734 }
735 }
736 }
737 };
738
739 u64::try_from(id).map_err(|_| {
740 MutationExecutorError::Bind(format!("generated id {id} cannot be represented as u64"))
741 })
742 }
743}
744
745impl InternalIdGenerator for PgIdSpaceGenerator {
746 fn generate_id(&self, entity: &str) -> Result<u64, RuntimeError> {
747 let generator = self.clone();
748 let entity = entity.to_owned();
749 block_on_id_generation(async move { generator.next_id(&entity).await })
750 }
751}
752
753fn block_on_id_generation<F>(future: F) -> Result<u64, RuntimeError>
754where
755 F: Future<Output = Result<u64, MutationExecutorError>> + Send + 'static,
756{
757 let result = match tokio::runtime::Handle::try_current() {
758 Ok(handle) => tokio::task::block_in_place(|| handle.block_on(future)),
759 Err(_) => tokio::runtime::Builder::new_current_thread()
760 .enable_all()
761 .build()
762 .map_err(|err| RuntimeError::IdGeneration(err.to_string()))?
763 .block_on(future),
764 };
765 result.map_err(|err| RuntimeError::IdGeneration(err.to_string()))
766}
767
768fn quote_ident(ident: &str) -> String {
769 quote_identifier_if_needed(ident, '"')
770}
771
772fn strip_identifier_quotes(ident: &str) -> &str {
776 let bytes = ident.as_bytes();
777 if bytes.len() >= 2 {
778 let (first, last) = (bytes[0], bytes[bytes.len() - 1]);
779 if (first == b'"' && last == b'"')
780 || (first == b'`' && last == b'`')
781 || (first == b'[' && last == b']')
782 {
783 return &ident[1..ident.len() - 1];
784 }
785 }
786 ident
787}
788
789fn try_parse_datetime_from_str(s: &str) -> Option<chrono::DateTime<chrono::Utc>> {
790 if let Ok(dt) = chrono::DateTime::parse_from_rfc3339(s) {
791 return Some(dt.with_timezone(&chrono::Utc));
792 }
793 if let Ok(ndt) = chrono::NaiveDateTime::parse_from_str(s, "%Y-%m-%d %H:%M:%S") {
794 return Some(chrono::DateTime::from_naive_utc_and_offset(
795 ndt,
796 chrono::Utc,
797 ));
798 }
799 if let Ok(nd) = chrono::NaiveDate::parse_from_str(s, "%Y-%m-%d") {
800 let ndt = nd.and_hms_opt(0, 0, 0)?;
801 return Some(chrono::DateTime::from_naive_utc_and_offset(
802 ndt,
803 chrono::Utc,
804 ));
805 }
806 None
807}
808
809#[derive(Debug, Clone, Copy)]
810struct PgNull;
811
812impl tokio_postgres::types::ToSql for PgNull {
813 fn to_sql(
814 &self,
815 ty: &tokio_postgres::types::Type,
816 out: &mut bytes::BytesMut,
817 ) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>> {
818 Ok(tokio_postgres::types::IsNull::Yes)
819 }
820
821 fn accepts(ty: &tokio_postgres::types::Type) -> bool {
822 true
823 }
824
825 fn to_sql_checked(
826 &self,
827 ty: &tokio_postgres::types::Type,
828 out: &mut bytes::BytesMut,
829 ) -> Result<tokio_postgres::types::IsNull, Box<dyn std::error::Error + Sync + Send>> {
830 Ok(tokio_postgres::types::IsNull::Yes)
831 }
832}
833
834struct PgArgs {
835 values: Vec<Box<dyn tokio_postgres::types::ToSql + Sync + Send>>,
836}
837impl PgArgs {
838 fn add<T: tokio_postgres::types::ToSql + Sync + Send + 'static>(&mut self, v: T) {
839 self.values.push(Box::new(v));
840 }
841 fn as_refs(&self) -> Vec<&(dyn tokio_postgres::types::ToSql + Sync)> {
842 self.values.iter().map(|b| b.as_ref() as _).collect()
843 }
844}
845
846fn bind_pg(args: &mut PgArgs, value: &Value) -> Result<(), MutationExecutorError> {
847 match value {
848 Value::Null => {
849 args.add(PgNull);
850 }
851 Value::Bool(v) => args.add(*v),
852 Value::I64(v) => args.add(*v),
853 Value::U64(v) => {
854 let v = i64::try_from(*v).map_err(|_| {
855 MutationExecutorError::Bind(format!("u64 value {v} exceeds i64 range"))
856 })?;
857 args.add(v);
858 }
859 Value::F64(v) => args.add(*v),
860 Value::Decimal(v) => args.add(*v),
861 Value::Text(v) => match try_parse_datetime_from_str(v) {
862 Some(dt) => args.add(dt),
863 None => args.add(v.clone()),
864 },
865 Value::Json(v) => {
866 let j_val: serde_json::Value =
867 serde_json::to_value(v).map_err(|e| MutationExecutorError::Bind(e.to_string()))?;
868 args.add(j_val);
869 }
870 Value::Date(v) => args.add(*v),
871 Value::Timestamp(v) => args.add(v.to_datetime()),
872 Value::Object(_) => return Err(MutationExecutorError::UnsupportedValue("object")),
873 Value::List(values) => bind_pg_list(args, values)?,
874 Value::TypedNull(dt) => match dt {
875 DataType::Bool => args.add(Option::<bool>::None),
876 DataType::I64 | DataType::U64 => args.add(Option::<i64>::None),
877 DataType::F64 => args.add(Option::<f64>::None),
878 DataType::Decimal => args.add(Option::<Decimal>::None),
879 DataType::Text | DataType::LargeText => args.add(Option::<String>::None),
880 DataType::Json => args.add(Option::<serde_json::Value>::None),
881 DataType::Date => args.add(Option::<NaiveDate>::None),
882 DataType::Timestamp => args.add(Option::<DateTime<Utc>>::None),
883 },
884 }
885 Ok(())
886}
887
888fn bind_pg_list(args: &mut PgArgs, values: &[Value]) -> Result<(), MutationExecutorError> {
889 let Some(first) = values.first() else {
890 return Err(MutationExecutorError::UnsupportedValue("empty list"));
891 };
892 match first {
893 Value::Bool(_) => {
894 let values = values
895 .iter()
896 .map(|value| match value {
897 Value::Bool(value) => Ok(*value),
898 _ => Err(MutationExecutorError::UnsupportedValue("mixed bool list")),
899 })
900 .collect::<Result<Vec<_>, _>>()?;
901 args.add(values);
902 }
903 Value::I64(_) => {
904 let values = values
905 .iter()
906 .map(|value| match value {
907 Value::I64(value) => Ok(*value),
908 _ => Err(MutationExecutorError::UnsupportedValue("mixed i64 list")),
909 })
910 .collect::<Result<Vec<_>, _>>()?;
911 args.add(values);
912 }
913 Value::U64(_) => {
914 let values = values
915 .iter()
916 .map(|value| match value {
917 Value::U64(value) => i64::try_from(*value).map_err(|_| {
918 MutationExecutorError::Bind(format!("u64 value {value} exceeds i64 range"))
919 }),
920 _ => Err(MutationExecutorError::UnsupportedValue("mixed u64 list")),
921 })
922 .collect::<Result<Vec<_>, _>>()?;
923 args.add(values);
924 }
925 Value::F64(_) => {
926 let values = values
927 .iter()
928 .map(|value| match value {
929 Value::F64(value) => Ok(*value),
930 _ => Err(MutationExecutorError::UnsupportedValue("mixed f64 list")),
931 })
932 .collect::<Result<Vec<_>, _>>()?;
933 args.add(values);
934 }
935 Value::Decimal(_) => {
936 let values = values
937 .iter()
938 .map(|value| match value {
939 Value::Decimal(value) => Ok(*value),
940 _ => Err(MutationExecutorError::UnsupportedValue(
941 "mixed decimal list",
942 )),
943 })
944 .collect::<Result<Vec<_>, _>>()?;
945 args.add(values);
946 }
947 Value::Text(_) => {
948 let values = values
949 .iter()
950 .map(|value| match value {
951 Value::Text(value) => Ok(value.clone()),
952 _ => Err(MutationExecutorError::UnsupportedValue("mixed text list")),
953 })
954 .collect::<Result<Vec<_>, _>>()?;
955 args.add(values);
956 }
957 Value::Date(_) => {
958 let values = values
959 .iter()
960 .map(|value| match value {
961 Value::Date(value) => Ok(*value),
962 _ => Err(MutationExecutorError::UnsupportedValue("mixed date list")),
963 })
964 .collect::<Result<Vec<_>, _>>()?;
965 args.add(values);
966 }
967 Value::Timestamp(_) => {
968 let values = values
969 .iter()
970 .map(|value| match value {
971 Value::Timestamp(value) => Ok(value.to_datetime()),
972 _ => Err(MutationExecutorError::UnsupportedValue(
973 "mixed timestamp list",
974 )),
975 })
976 .collect::<Result<Vec<_>, _>>()?;
977 args.add(values);
978 }
979 Value::Null => return Err(MutationExecutorError::UnsupportedValue("null list")),
980 Value::Json(_) => return Err(MutationExecutorError::UnsupportedValue("json list")),
981 Value::Object(_) => return Err(MutationExecutorError::UnsupportedValue("object list")),
982 Value::List(_) => return Err(MutationExecutorError::UnsupportedValue("nested list")),
983 Value::TypedNull(_) => return Err(MutationExecutorError::UnsupportedValue("null list")),
984 }
985 Ok(())
986}
987
988fn decode_pg_row(row: &tokio_postgres::Row) -> Result<Record, MutationExecutorError> {
989 let mut record = BTreeMap::new();
990 for (index, column) in row.columns().iter().enumerate() {
991 let name = column.name().to_owned();
992 let type_name = column.type_().name().to_ascii_uppercase();
993
994 let value = match type_name.as_str() {
995 "BOOL" | "BOOLEAN" => {
996 let v: Option<bool> = row.try_get(index)?;
997 match v {
998 Some(v) => Value::Bool(v),
999 None => Value::Null,
1000 }
1001 }
1002 "INT2" => {
1003 let v: Option<i16> = row.try_get(index)?;
1004 match v {
1005 Some(v) => Value::I64(v as i64),
1006 None => Value::Null,
1007 }
1008 }
1009 "INT4" => {
1010 let v: Option<i32> = row.try_get(index)?;
1011 match v {
1012 Some(v) => Value::I64(v as i64),
1013 None => Value::Null,
1014 }
1015 }
1016 "INT8" => {
1017 let v: Option<i64> = row.try_get(index)?;
1018 match v {
1019 Some(v) => Value::I64(v),
1020 None => Value::Null,
1021 }
1022 }
1023 "FLOAT4" => {
1024 let v: Option<f32> = row.try_get(index)?;
1025 match v {
1026 Some(v) => Value::F64(v as f64),
1027 None => Value::Null,
1028 }
1029 }
1030 "FLOAT8" => {
1031 let v: Option<f64> = row.try_get(index)?;
1032 match v {
1033 Some(v) => Value::F64(v),
1034 None => Value::Null,
1035 }
1036 }
1037 "NUMERIC" => {
1038 let v: Option<Decimal> = row.try_get(index)?;
1039 match v {
1040 Some(v) => Value::Decimal(v),
1041 None => Value::Null,
1042 }
1043 }
1044 "JSON" | "JSONB" => {
1045 let v: Option<serde_json::Value> = row.try_get(index)?;
1046 match v {
1047 Some(j) => Value::Json(j.into()),
1048 None => Value::Null,
1049 }
1050 }
1051 "DATE" => {
1052 let v: Option<NaiveDate> = row.try_get(index)?;
1053 match v {
1054 Some(v) => Value::Date(v),
1055 None => Value::Null,
1056 }
1057 }
1058 "TIMESTAMP" | "TIMESTAMPTZ" => {
1059 let v: Option<DateTime<Utc>> = row.try_get(index)?;
1060 match v {
1061 Some(v) => Value::Timestamp(teaql_core::time::Timestamp(v.timestamp_millis())),
1062 None => Value::Null,
1063 }
1064 }
1065 "TEXT" | "VARCHAR" | "BPCHAR" | "NAME" | "UUID" => {
1066 let v: Option<String> = row.try_get(index)?;
1067 match v {
1068 Some(v) => Value::Text(v),
1069 None => Value::Null,
1070 }
1071 }
1072 other => {
1073 return Err(MutationExecutorError::UnsupportedColumnType(
1074 other.to_owned(),
1075 ));
1076 }
1077 };
1078 record.insert(name, value);
1079 }
1080 Ok(record)
1081}
1082
1083#[cfg(test)]
1084mod tests {
1085 use super::*;
1086 use teaql_core::{DeleteCommand, RecoverCommand};
1087
1088 fn entity() -> EntityDescriptor {
1089 EntityDescriptor::new("Order")
1090 .table_name("orders")
1091 .property(
1092 PropertyDescriptor::new("id", DataType::U64)
1093 .column_name("id")
1094 .id()
1095 .not_null(),
1096 )
1097 .property(
1098 PropertyDescriptor::new("version", DataType::I64)
1099 .column_name("version")
1100 .version()
1101 .not_null(),
1102 )
1103 .property(PropertyDescriptor::new("name", DataType::Text).column_name("name"))
1104 }
1105
1106 #[test]
1107 fn postgres_dialect_compiles_mutations_with_numbered_placeholders() {
1108 let insert = PostgresDialect
1109 .compile_insert(
1110 &entity(),
1111 &InsertCommand::new("Order")
1112 .value("id", 1_u64)
1113 .value("name", "A"),
1114 )
1115 .unwrap();
1116 assert_eq!(insert.sql, "INSERT INTO orders (id, name) VALUES ($1, $2)");
1117
1118 let update = PostgresDialect
1119 .compile_update(
1120 &entity(),
1121 &UpdateCommand::new("Order", 1_u64)
1122 .expected_version(3)
1123 .value("name", "B"),
1124 )
1125 .unwrap();
1126 assert_eq!(
1127 update.sql,
1128 "UPDATE orders SET name = $1, version = $2 WHERE id = $3 AND version = $4"
1129 );
1130
1131 let delete = PostgresDialect
1132 .compile_delete(
1133 &entity(),
1134 &DeleteCommand::new("Order", 1_u64).expected_version(3),
1135 )
1136 .unwrap();
1137 let recover = PostgresDialect
1138 .compile_recover(&entity(), &RecoverCommand::new("Order", 1_u64, -4))
1139 .unwrap();
1140 assert_eq!(
1141 delete.sql,
1142 "UPDATE orders SET version = $1 WHERE id = $2 AND version = $3"
1143 );
1144 assert_eq!(
1145 recover.sql,
1146 "UPDATE orders SET version = $1 WHERE id = $2 AND version = $3"
1147 );
1148 }
1149
1150 #[test]
1151 fn postgres_dialect_compiles_schema_and_large_in_array_binds() {
1152 let create = PostgresDialect.compile_create_table(&entity()).unwrap();
1153 assert_eq!(
1154 create,
1155 "CREATE TABLE IF NOT EXISTS orders (id BIGINT PRIMARY KEY NOT NULL, version BIGINT NOT NULL, name VARCHAR(255))"
1156 );
1157 assert!(
1158 PostgresDialect
1159 .schema_setup_sqls()
1160 .iter()
1161 .any(|sql| sql.contains("CREATE OR REPLACE FUNCTION soundex"))
1162 );
1163
1164 let query = PostgresDialect
1165 .compile_select(
1166 &entity(),
1167 &SelectQuery::new("Order")
1168 .filter(Expr::in_large(
1169 "id",
1170 vec![Value::from(1_u64), Value::from(2_u64)],
1171 ))
1172 .order_asc("id"),
1173 )
1174 .unwrap();
1175 assert_eq!(
1176 query.sql,
1177 "SELECT id, version, name FROM orders WHERE (id = ANY($1)) ORDER BY id ASC"
1178 );
1179 assert_eq!(
1180 query.params,
1181 vec![Value::List(vec![Value::U64(1), Value::U64(2)])]
1182 );
1183 }
1184}