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