Skip to main content

teaql_provider_postgres/
lib.rs

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