Skip to main content

risingwave_connector/sink/
postgres.rs

1// Copyright 2024 RisingWave Labs
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::collections::{BTreeMap, HashSet};
16use std::sync::Arc;
17
18use anyhow::{Context, anyhow};
19use async_trait::async_trait;
20use futures::StreamExt;
21use futures::stream::FuturesUnordered;
22use itertools::Itertools;
23use phf::phf_set;
24use risingwave_common::array::{Op, StreamChunk};
25use risingwave_common::catalog::Schema;
26use risingwave_common::row::{Row, RowExt};
27use serde::Deserialize;
28use serde_with::{DisplayFromStr, serde_as};
29use simd_json::prelude::ArrayTrait;
30use thiserror_ext::AsReport;
31use tokio_postgres::types::Type as PgType;
32
33use super::{
34    LogSinker, SINK_TYPE_APPEND_ONLY, SINK_TYPE_OPTION, SINK_TYPE_UPSERT, SinkError, SinkLogReader,
35};
36use crate::connector_common::{
37    PgConnectionConfig, PostgresExternalTable, SslMode, TcpKeepaliveConfig, create_pg_client,
38};
39use crate::enforce_secret::EnforceSecret;
40use crate::parser::scalar_adapter::{ScalarAdapter, validate_pg_type_to_rw_type};
41use crate::sink::log_store::{LogStoreReadItem, TruncateOffset};
42use crate::sink::{Result, Sink, SinkParam, SinkWriterParam};
43
44pub const POSTGRES_SINK: &str = "postgres";
45
46const CHECK_FOREIGN_KEY_SQL: &str = r#"
47    SELECT EXISTS (
48        SELECT 1
49        FROM pg_constraint c
50        JOIN pg_class t ON t.oid = c.conrelid
51        JOIN pg_namespace n ON n.oid = t.relnamespace
52        WHERE n.nspname = $1
53          AND t.relname = $2
54          AND c.contype = 'f'
55    )
56"#;
57
58#[serde_as]
59#[derive(Clone, Debug, Deserialize)]
60pub struct PostgresConfig {
61    pub host: String,
62    #[serde_as(as = "DisplayFromStr")]
63    pub port: u16,
64    pub user: String,
65    pub password: String,
66    pub database: String,
67    pub table: String,
68    #[serde(default = "default_schema")]
69    pub schema: String,
70    #[serde(default = "Default::default")]
71    pub ssl_mode: SslMode,
72    #[serde(rename = "ssl.root.cert")]
73    pub ssl_root_cert: Option<String>,
74    #[serde(default = "default_max_batch_rows")]
75    #[serde_as(as = "DisplayFromStr")]
76    pub max_batch_rows: usize,
77    pub r#type: String, // accept "append-only" or "upsert"
78    #[serde(default, rename = "tcp.keepalive.enable")]
79    #[serde_as(as = "DisplayFromStr")]
80    pub tcp_keepalive_enable: bool,
81
82    #[serde(flatten)]
83    pub tcp_keepalive: Option<TcpKeepaliveConfig>,
84
85    #[serde(flatten)]
86    pub unknown_fields: std::collections::HashMap<String, String>,
87}
88
89crate::impl_sink_unknown_fields!(PostgresConfig);
90
91impl EnforceSecret for PostgresConfig {
92    const ENFORCE_SECRET_PROPERTIES: phf::Set<&'static str> = phf_set! {
93        "password", "ssl.root.cert"
94    };
95}
96
97fn default_max_batch_rows() -> usize {
98    1024
99}
100
101fn default_schema() -> String {
102    "public".to_owned()
103}
104
105fn tcp_keepalive_from_config(config: &PostgresConfig) -> Option<TcpKeepaliveConfig> {
106    if config.tcp_keepalive_enable {
107        config
108            .tcp_keepalive
109            .clone()
110            .or_else(|| Some(TcpKeepaliveConfig::default()))
111    } else {
112        None
113    }
114}
115
116async fn ensure_no_foreign_key(config: &PostgresConfig) -> Result<()> {
117    let pg_conn = config.pg_connection_config();
118    let client = create_pg_client(&pg_conn, tcp_keepalive_from_config(config)).await?;
119
120    ensure_no_foreign_key_with_client(&client, &config.schema, &config.table).await
121}
122
123async fn ensure_no_foreign_key_with_client(
124    client: &tokio_postgres::Client,
125    schema: &str,
126    table: &str,
127) -> Result<()> {
128    let has_foreign_key = client
129        .query_one(CHECK_FOREIGN_KEY_SQL, &[&schema, &table])
130        .await
131        .context("failed to check foreign key constraints")?
132        .get::<_, bool>(0);
133
134    if has_foreign_key {
135        return Err(SinkError::Config(anyhow!(
136            "Postgres sink does not support target table \"{}\".\"{}\" with foreign key constraints. Please remove foreign key constraints from the target table or choose a different sink table.",
137            schema,
138            table,
139        )));
140    }
141
142    Ok(())
143}
144
145impl PostgresConfig {
146    pub fn from_btreemap(properties: BTreeMap<String, String>) -> Result<Self> {
147        let config =
148            serde_json::from_value::<PostgresConfig>(serde_json::to_value(properties).unwrap())
149                .map_err(|e| SinkError::Config(anyhow!(e)))?;
150        if config.r#type != SINK_TYPE_APPEND_ONLY && config.r#type != SINK_TYPE_UPSERT {
151            return Err(SinkError::Config(anyhow!(
152                "`{}` must be {}, or {}",
153                SINK_TYPE_OPTION,
154                SINK_TYPE_APPEND_ONLY,
155                SINK_TYPE_UPSERT
156            )));
157        }
158        Ok(config)
159    }
160
161    pub fn pg_connection_config(&self) -> PgConnectionConfig {
162        PgConnectionConfig {
163            host: self.host.clone(),
164            port: self.port,
165            user: self.user.clone(),
166            password: self.password.clone(),
167            database: self.database.clone(),
168            ssl_mode: self.ssl_mode.clone(),
169            ssl_root_cert: self.ssl_root_cert.clone(),
170        }
171    }
172}
173
174#[derive(Debug)]
175pub struct PostgresSink {
176    pub config: PostgresConfig,
177    schema: Schema,
178    pk_indices: Vec<usize>,
179    is_append_only: bool,
180}
181
182impl PostgresSink {
183    pub fn new(
184        config: PostgresConfig,
185        schema: Schema,
186        pk_indices: Vec<usize>,
187        is_append_only: bool,
188    ) -> Result<Self> {
189        Ok(Self {
190            config,
191            schema,
192            pk_indices,
193            is_append_only,
194        })
195    }
196}
197
198impl EnforceSecret for PostgresSink {
199    fn enforce_secret<'a>(
200        prop_iter: impl Iterator<Item = &'a str>,
201    ) -> crate::error::ConnectorResult<()> {
202        for prop in prop_iter {
203            PostgresConfig::enforce_one(prop)?;
204        }
205        Ok(())
206    }
207}
208
209impl TryFrom<SinkParam> for PostgresSink {
210    type Error = SinkError;
211
212    fn try_from(param: SinkParam) -> std::result::Result<Self, Self::Error> {
213        let schema = param.schema();
214        let pk_indices = param.downstream_pk_or_empty();
215        let config = PostgresConfig::from_btreemap(param.properties)?;
216        PostgresSink::new(config, schema, pk_indices, param.sink_type.is_append_only())
217    }
218}
219
220impl Sink for PostgresSink {
221    type LogSinker = PostgresSinkWriter;
222
223    const SINK_NAME: &'static str = POSTGRES_SINK;
224
225    crate::impl_validate_sink_unknown_fields!();
226
227    async fn validate(&self) -> Result<()> {
228        if !self.is_append_only && self.pk_indices.is_empty() {
229            return Err(SinkError::Config(anyhow!(
230                "Primary key not defined for upsert Postgres sink (please define in `primary_key` field)"
231            )));
232        }
233
234        ensure_no_foreign_key(&self.config).await?;
235
236        // Verify our sink schema is compatible with Postgres
237        {
238            let pg_conn = self.config.pg_connection_config();
239            let pg_table = PostgresExternalTable::connect(
240                &pg_conn,
241                &self.config.schema,
242                &self.config.table,
243                self.is_append_only,
244                None,
245            )
246            .await
247            .context(format!(
248                "failed to connect to database: {}, schema: {}, table: {}",
249                self.config.database, self.config.schema, self.config.table
250            ))?;
251
252            // Check that names and types match, order of columns doesn't matter.
253            {
254                let pg_columns = pg_table.column_descs();
255                let sink_columns = self.schema.fields();
256                if pg_columns.len() < sink_columns.len() {
257                    return Err(SinkError::Config(anyhow!(
258                        "Column count mismatch: Postgres table has {} columns, but sink schema has {} columns, sink should have less or equal columns to the Postgres table",
259                        pg_columns.len(),
260                        sink_columns.len()
261                    )));
262                }
263
264                let pg_columns_lookup = pg_columns
265                    .iter()
266                    .map(|c| (c.name.clone(), c.data_type.clone()))
267                    .collect::<BTreeMap<_, _>>();
268                for sink_column in sink_columns {
269                    let pg_column = pg_columns_lookup.get(&sink_column.name);
270                    match pg_column {
271                        None => {
272                            return Err(SinkError::Config(anyhow!(
273                                "Column `{}` not found in Postgres table `{}`",
274                                sink_column.name,
275                                self.config.table
276                            )));
277                        }
278                        Some(pg_column) => {
279                            if !validate_pg_type_to_rw_type(pg_column, &sink_column.data_type()) {
280                                return Err(SinkError::Config(anyhow!(
281                                    "Column `{}` in Postgres table `{}` has type `{}`, but sink schema defines it as type `{}`",
282                                    sink_column.name,
283                                    self.config.table,
284                                    pg_column,
285                                    sink_column.data_type()
286                                )));
287                            }
288                        }
289                    }
290                }
291            }
292
293            // check that pk matches
294            {
295                let pg_pk_names = pg_table.pk_names();
296                let sink_pk_names = self
297                    .pk_indices
298                    .iter()
299                    .map(|i| &self.schema.fields()[*i].name)
300                    .collect::<HashSet<_>>();
301                if pg_pk_names.len() != sink_pk_names.len() {
302                    return Err(SinkError::Config(anyhow!(
303                        "Primary key mismatch: Postgres table has primary key on columns {:?}, but sink schema defines primary key on columns {:?}",
304                        pg_pk_names,
305                        sink_pk_names
306                    )));
307                }
308                for name in pg_pk_names {
309                    if !sink_pk_names.contains(name) {
310                        return Err(SinkError::Config(anyhow!(
311                            "Primary key mismatch: Postgres table has primary key on column `{}`, but sink schema does not define it as a primary key",
312                            name
313                        )));
314                    }
315                }
316            }
317        }
318
319        Ok(())
320    }
321
322    async fn new_log_sinker(&self, _writer_param: SinkWriterParam) -> Result<Self::LogSinker> {
323        PostgresSinkWriter::new(
324            self.config.clone(),
325            self.schema.clone(),
326            self.pk_indices.clone(),
327            self.is_append_only,
328        )
329        .await
330    }
331}
332
333pub struct PostgresSinkWriter {
334    is_append_only: bool,
335    client: tokio_postgres::Client,
336    pk_indices: Vec<usize>,
337    pk_types: Vec<PgType>,
338    schema_types: Vec<PgType>,
339    raw_insert_sql: Arc<String>,
340    raw_upsert_sql: Arc<String>,
341    raw_delete_sql: Arc<String>,
342    insert_sql: Arc<tokio_postgres::Statement>,
343    delete_sql: Arc<tokio_postgres::Statement>,
344    upsert_sql: Arc<tokio_postgres::Statement>,
345}
346
347impl PostgresSinkWriter {
348    async fn new(
349        config: PostgresConfig,
350        schema: Schema,
351        pk_indices: Vec<usize>,
352        is_append_only: bool,
353    ) -> Result<Self> {
354        let tcp_keepalive = tcp_keepalive_from_config(&config);
355
356        let pg_conn = config.pg_connection_config();
357        let client = create_pg_client(&pg_conn, tcp_keepalive).await?;
358
359        ensure_no_foreign_key_with_client(&client, &config.schema, &config.table).await?;
360
361        let pk_indices_lookup = pk_indices.iter().copied().collect::<HashSet<_>>();
362
363        // Rewrite schema types for serialization
364        let (pk_types, schema_types) = {
365            let name_to_type = PostgresExternalTable::type_mapping(
366                &pg_conn,
367                &config.schema,
368                &config.table,
369                is_append_only,
370            )
371            .await?;
372            let mut schema_types = Vec::with_capacity(schema.fields.len());
373            let mut pk_types = Vec::with_capacity(pk_indices.len());
374            for (i, field) in schema.fields.iter().enumerate() {
375                let field_name = &field.name;
376                let actual_data_type = name_to_type.get(field_name).map(|t| (*t).clone());
377                let actual_data_type = actual_data_type
378                    .ok_or_else(|| {
379                        SinkError::Config(anyhow!(
380                            "Column `{}` not found in sink schema",
381                            field_name
382                        ))
383                    })?
384                    .clone();
385                if pk_indices_lookup.contains(&i) {
386                    pk_types.push(actual_data_type.clone())
387                }
388                schema_types.push(actual_data_type);
389            }
390            (pk_types, schema_types)
391        };
392
393        let raw_insert_sql = create_insert_sql(&schema, &config.schema, &config.table);
394        let raw_upsert_sql = create_upsert_sql(
395            &schema,
396            &config.schema,
397            &config.table,
398            &pk_indices,
399            &pk_indices_lookup,
400        );
401        let raw_delete_sql = create_delete_sql(&schema, &config.schema, &config.table, &pk_indices);
402
403        let insert_sql = client
404            .prepare(&raw_insert_sql)
405            .await
406            .with_context(|| format!("failed to prepare insert statement: {}", raw_insert_sql))?;
407        let upsert_sql = client
408            .prepare(&raw_upsert_sql)
409            .await
410            .with_context(|| format!("failed to prepare upsert statement: {}", raw_upsert_sql))?;
411        let delete_sql = client
412            .prepare(&raw_delete_sql)
413            .await
414            .with_context(|| format!("failed to prepare delete statement: {}", raw_delete_sql))?;
415
416        let writer = Self {
417            is_append_only,
418            client,
419            pk_indices,
420            pk_types,
421            schema_types,
422            raw_insert_sql: Arc::new(raw_insert_sql),
423            raw_upsert_sql: Arc::new(raw_upsert_sql),
424            raw_delete_sql: Arc::new(raw_delete_sql),
425            insert_sql: Arc::new(insert_sql),
426            delete_sql: Arc::new(delete_sql),
427            upsert_sql: Arc::new(upsert_sql),
428        };
429        Ok(writer)
430    }
431
432    async fn write_batch(&mut self, chunk: StreamChunk) -> Result<()> {
433        // https://www.postgresql.org/docs/current/limits.html
434        // We have a limit of 65,535 parameters in a single query, as restricted by the PostgreSQL protocol.
435        if self.is_append_only {
436            self.write_batch_append_only(chunk).await
437        } else {
438            self.write_batch_non_append_only(chunk).await
439        }
440    }
441
442    async fn write_batch_append_only(&mut self, chunk: StreamChunk) -> Result<()> {
443        let transaction = Arc::new(self.client.transaction().await?);
444        let mut insert_futures = FuturesUnordered::new();
445        for (op, row) in chunk.rows() {
446            match op {
447                Op::Insert => {
448                    let pg_row = convert_row_to_pg_row(row, &self.schema_types);
449                    let insert_sql = self.insert_sql.clone();
450                    let raw_insert_sql = self.raw_insert_sql.clone();
451                    let transaction = transaction.clone();
452                    let future = async move {
453                        transaction
454                            .execute_raw(insert_sql.as_ref(), &pg_row)
455                            .await
456                            .with_context(|| {
457                                format!(
458                                    "failed to execute insert statement: {}, parameters: {:?}",
459                                    raw_insert_sql, pg_row
460                                )
461                            })
462                    };
463                    insert_futures.push(future);
464                }
465                _ => {
466                    tracing::error!(
467                        "row ignored, append-only sink should not receive update insert, update delete and delete operations"
468                    );
469                }
470            }
471        }
472
473        while let Some(result) = insert_futures.next().await {
474            result?;
475        }
476        if let Some(transaction) = Arc::into_inner(transaction) {
477            transaction.commit().await?;
478        } else {
479            tracing::error!("transaction lost!");
480        }
481
482        Ok(())
483    }
484
485    async fn write_batch_non_append_only(&mut self, chunk: StreamChunk) -> Result<()> {
486        let transaction = Arc::new(self.client.transaction().await?);
487        let mut delete_futures = FuturesUnordered::new();
488        let mut upsert_futures = FuturesUnordered::new();
489        for (op, row) in chunk.rows() {
490            match op {
491                Op::Delete | Op::UpdateDelete => {
492                    let pg_row =
493                        convert_row_to_pg_row(row.project(&self.pk_indices), &self.pk_types);
494                    let delete_sql = self.delete_sql.clone();
495                    let raw_delete_sql = self.raw_delete_sql.clone();
496                    let transaction = transaction.clone();
497                    let future = async move {
498                        transaction
499                            .execute_raw(delete_sql.as_ref(), &pg_row)
500                            .await
501                            .with_context(|| {
502                                format!(
503                                    "failed to execute delete statement: {}, parameters: {:?}",
504                                    raw_delete_sql, pg_row
505                                )
506                            })
507                    };
508                    delete_futures.push(future);
509                }
510                Op::Insert | Op::UpdateInsert => {
511                    let pg_row = convert_row_to_pg_row(row, &self.schema_types);
512                    let upsert_sql = self.upsert_sql.clone();
513                    let raw_upsert_sql = self.raw_upsert_sql.clone();
514                    let transaction = transaction.clone();
515                    let future = async move {
516                        transaction
517                            .execute_raw(upsert_sql.as_ref(), &pg_row)
518                            .await
519                            .with_context(|| {
520                                format!(
521                                    "failed to execute upsert statement: {}, parameters: {:?}",
522                                    raw_upsert_sql, pg_row
523                                )
524                            })
525                    };
526                    upsert_futures.push(future);
527                }
528            }
529        }
530        while let Some(result) = delete_futures.next().await {
531            result?;
532        }
533        while let Some(result) = upsert_futures.next().await {
534            result?;
535        }
536        if let Some(transaction) = Arc::into_inner(transaction) {
537            transaction.commit().await?;
538        } else {
539            tracing::error!("transaction lost!");
540        }
541        Ok(())
542    }
543}
544
545#[async_trait]
546impl LogSinker for PostgresSinkWriter {
547    async fn consume_log_and_sink(mut self, mut log_reader: impl SinkLogReader) -> Result<!> {
548        log_reader.start_from(None).await?;
549        loop {
550            let (epoch, item) = log_reader.next_item().await?;
551            match item {
552                LogStoreReadItem::StreamChunk { chunk, chunk_id } => {
553                    self.write_batch(chunk).await?;
554                    log_reader.truncate(TruncateOffset::Chunk { epoch, chunk_id })?;
555                }
556                LogStoreReadItem::Barrier { .. } => {
557                    log_reader.truncate(TruncateOffset::Barrier { epoch })?;
558                }
559            }
560        }
561    }
562}
563
564fn create_insert_sql(schema: &Schema, schema_name: &str, table_name: &str) -> String {
565    let normalized_table_name = format!(
566        "{}.{}",
567        quote_identifier(schema_name),
568        quote_identifier(table_name)
569    );
570    let number_of_columns = schema.len();
571    let columns: String = schema
572        .fields()
573        .iter()
574        .map(|field| quote_identifier(&field.name))
575        .join(", ");
576    let column_parameters: String = (0..number_of_columns)
577        .map(|i| format!("${}", i + 1))
578        .join(", ");
579    format!("INSERT INTO {normalized_table_name} ({columns}) VALUES ({column_parameters})")
580}
581
582fn create_delete_sql(
583    schema: &Schema,
584    schema_name: &str,
585    table_name: &str,
586    pk_indices: &[usize],
587) -> String {
588    let normalized_table_name = format!(
589        "{}.{}",
590        quote_identifier(schema_name),
591        quote_identifier(table_name)
592    );
593    let pk_indices = if pk_indices.is_empty() {
594        (0..schema.len()).collect_vec()
595    } else {
596        pk_indices.to_vec()
597    };
598    let pk = {
599        let pk_symbols = pk_indices
600            .iter()
601            .map(|pk_index| quote_identifier(&schema.fields()[*pk_index].name))
602            .join(", ");
603        format!("({})", pk_symbols)
604    };
605    let parameters: String = (0..pk_indices.len())
606        .map(|i| format!("${}", i + 1))
607        .join(", ");
608    format!("DELETE FROM {normalized_table_name} WHERE {pk} in (({parameters}))")
609}
610
611fn create_upsert_sql(
612    schema: &Schema,
613    schema_name: &str,
614    table_name: &str,
615    pk_indices: &[usize],
616    pk_indices_lookup: &HashSet<usize>,
617) -> String {
618    let insert_sql = create_insert_sql(schema, schema_name, table_name);
619    if pk_indices.is_empty() {
620        return insert_sql;
621    }
622    let pk_columns = pk_indices
623        .iter()
624        .map(|pk_index| quote_identifier(&schema.fields()[*pk_index].name))
625        .collect_vec()
626        .join(", ");
627    let update_parameters: String = (0..schema.len())
628        .filter(|i| !pk_indices_lookup.contains(i))
629        .map(|i| {
630            let column = quote_identifier(&schema.fields()[i].name);
631            format!("{column} = EXCLUDED.{column}")
632        })
633        .collect_vec()
634        .join(", ");
635    if update_parameters.is_empty() {
636        format!("{insert_sql} on conflict ({pk_columns}) do nothing")
637    } else {
638        format!("{insert_sql} on conflict ({pk_columns}) do update set {update_parameters}")
639    }
640}
641
642/// Quote an identifier for PostgreSQL.
643fn quote_identifier(identifier: &str) -> String {
644    format!("\"{}\"", identifier.replace("\"", "\"\""))
645}
646
647type PgDatum = Option<ScalarAdapter>;
648type PgRow = Vec<PgDatum>;
649
650fn convert_row_to_pg_row(row: impl Row, schema_types: &[PgType]) -> PgRow {
651    let mut buffer = Vec::with_capacity(row.len());
652    for (i, datum_ref) in row.iter().enumerate() {
653        let pg_datum = datum_ref.map(|s| {
654            match ScalarAdapter::from_scalar(s, &schema_types[i]) {
655                Ok(scalar) => Some(scalar),
656                Err(e) => {
657                    tracing::error!(error=%e.as_report(), scalar=?s, "Failed to convert scalar to pg value");
658                    None
659                }
660            }
661        });
662        buffer.push(pg_datum.flatten());
663    }
664    buffer
665}
666
667#[cfg(test)]
668mod tests {
669    use std::fmt::Display;
670
671    use expect_test::{Expect, expect};
672    use risingwave_common::catalog::Field;
673    use risingwave_common::types::DataType;
674
675    use super::*;
676
677    fn check(actual: impl Display, expect: Expect) {
678        let actual = actual.to_string();
679        expect.assert_eq(&actual);
680    }
681
682    #[test]
683    fn test_create_insert_sql() {
684        let schema = Schema::new(vec![
685            Field {
686                data_type: DataType::Int32,
687                name: "a".to_owned(),
688            },
689            Field {
690                data_type: DataType::Int32,
691                name: "b".to_owned(),
692            },
693        ]);
694        let schema_name = "test_schema";
695        let table_name = "test_table";
696        let sql = create_insert_sql(&schema, schema_name, table_name);
697        check(
698            sql,
699            expect![[r#"INSERT INTO "test_schema"."test_table" ("a", "b") VALUES ($1, $2)"#]],
700        );
701    }
702
703    #[test]
704    fn test_create_delete_sql() {
705        let schema = Schema::new(vec![
706            Field {
707                data_type: DataType::Int32,
708                name: "a".to_owned(),
709            },
710            Field {
711                data_type: DataType::Int32,
712                name: "b".to_owned(),
713            },
714        ]);
715        let schema_name = "test_schema";
716        let table_name = "test_table";
717        let sql = create_delete_sql(&schema, schema_name, table_name, &[1]);
718        check(
719            sql,
720            expect![[r#"DELETE FROM "test_schema"."test_table" WHERE ("b") in (($1))"#]],
721        );
722        let table_name = "test_table";
723        let sql = create_delete_sql(&schema, schema_name, table_name, &[0, 1]);
724        check(
725            sql,
726            expect![[r#"DELETE FROM "test_schema"."test_table" WHERE ("a", "b") in (($1, $2))"#]],
727        );
728    }
729
730    #[test]
731    fn test_create_upsert_sql() {
732        let schema = Schema::new(vec![
733            Field {
734                data_type: DataType::Int32,
735                name: "a".to_owned(),
736            },
737            Field {
738                data_type: DataType::Int32,
739                name: "b".to_owned(),
740            },
741        ]);
742        let schema_name = "test_schema";
743        let table_name = "test_table";
744        let pk_indices_lookup = HashSet::from_iter([1]);
745        let sql = create_upsert_sql(&schema, schema_name, table_name, &[1], &pk_indices_lookup);
746        check(
747            sql,
748            expect![[
749                r#"INSERT INTO "test_schema"."test_table" ("a", "b") VALUES ($1, $2) on conflict ("b") do update set "a" = EXCLUDED."a""#
750            ]],
751        );
752    }
753
754    #[test]
755    fn test_create_upsert_sql_all_columns_are_primary_keys() {
756        let schema = Schema::new(vec![
757            Field {
758                data_type: DataType::Int32,
759                name: "user_id".to_owned(),
760            },
761            Field {
762                data_type: DataType::Int32,
763                name: "client_id".to_owned(),
764            },
765        ]);
766        let schema_name = "test_schema";
767        let table_name = "test_table";
768        let pk_indices_lookup = HashSet::from_iter([0, 1]);
769        let sql = create_upsert_sql(
770            &schema,
771            schema_name,
772            table_name,
773            &[0, 1],
774            &pk_indices_lookup,
775        );
776        check(
777            sql,
778            expect![[
779                r#"INSERT INTO "test_schema"."test_table" ("user_id", "client_id") VALUES ($1, $2) on conflict ("user_id", "client_id") do nothing"#
780            ]],
781        );
782    }
783}