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