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