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, 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
763/// Strip wrapping identifier quotes from a SQL identifier so that bare column
764/// names returned by `information_schema.columns` can be compared with
765/// potentially-quoted `PropertyDescriptor::column_name` values.
766fn 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}