Skip to main content

risingwave_connector/sink/snowflake_redshift/
snowflake.rs

1// Copyright 2025 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 core::num::NonZeroU64;
16use std::collections::BTreeMap;
17use std::time::Duration;
18
19use anyhow::anyhow;
20use phf::{Set, phf_set};
21use risingwave_common::array::StreamChunk;
22use risingwave_common::catalog::Schema;
23use risingwave_common::types::DataType;
24use risingwave_pb::connector_service::{SinkMetadata, sink_metadata};
25use risingwave_pb::stream_plan::PbSinkSchemaChange;
26use serde::Deserialize;
27use serde_with::{DisplayFromStr, serde_as};
28use thiserror_ext::AsReport;
29use tokio::sync::mpsc::{UnboundedSender, unbounded_channel};
30use tokio::time::{MissedTickBehavior, interval};
31use tonic::async_trait;
32use with_options::WithOptions;
33
34use crate::connector_common::IcebergSinkCompactionUpdate;
35use crate::enforce_secret::EnforceSecret;
36use crate::sink::catalog::SinkId;
37use crate::sink::coordinate::CoordinatedLogSinker;
38use crate::sink::decouple_checkpoint_log_sink::default_commit_checkpoint_interval;
39use crate::sink::file_sink::s3::S3Common;
40use crate::sink::jdbc_jni_client::{self, JdbcJniClient};
41use crate::sink::snowflake_redshift::{
42    __OP, __ROW_ID, SnowflakeRedshiftSinkJdbcWriter, SnowflakeRedshiftSinkS3Writer,
43};
44use crate::sink::writer::SinkWriter;
45use crate::sink::{
46    Result, SINK_TYPE_APPEND_ONLY, SINK_TYPE_OPTION, SINK_TYPE_UPSERT,
47    SinglePhaseCommitCoordinator, Sink, SinkCommitCoordinator, SinkError, SinkParam,
48    SinkWriterParam,
49};
50
51pub const SNOWFLAKE_SINK_V2: &str = "snowflake_v2";
52
53const AUTH_METHOD_PASSWORD: &str = "password";
54const AUTH_METHOD_KEY_PAIR_FILE: &str = "key_pair_file";
55const AUTH_METHOD_KEY_PAIR_OBJECT: &str = "key_pair_object";
56const PROP_AUTH_METHOD: &str = "auth.method";
57
58pub fn build_full_table_name(database: &str, schema_name: &str, table_name: &str) -> String {
59    format!(r#""{}"."{}"."{}""#, database, schema_name, table_name)
60}
61
62#[serde_as]
63#[derive(Debug, Clone, Deserialize, WithOptions)]
64pub struct SnowflakeV2Config {
65    #[serde(rename = "type")]
66    pub r#type: String,
67
68    #[serde(rename = "intermediate.table.name")]
69    pub snowflake_cdc_table_name: Option<String>,
70
71    #[serde(rename = "table.name")]
72    pub snowflake_target_table_name: Option<String>,
73
74    #[serde(rename = "database")]
75    pub snowflake_database: Option<String>,
76
77    #[serde(rename = "schema")]
78    pub snowflake_schema: Option<String>,
79
80    #[serde(default = "default_target_interval_schedule")]
81    #[serde(rename = "write.target.interval.seconds")]
82    #[serde_as(as = "DisplayFromStr")]
83    pub writer_target_interval_seconds: u64,
84
85    #[serde(default = "default_intermediate_interval_schedule")]
86    #[serde(rename = "write.intermediate.interval.seconds")]
87    #[serde_as(as = "DisplayFromStr")]
88    pub write_intermediate_interval_seconds: u64,
89
90    #[serde(rename = "warehouse")]
91    pub snowflake_warehouse: Option<String>,
92
93    #[serde(default, rename = "task.serverless")]
94    #[serde_as(as = "DisplayFromStr")]
95    pub task_serverless: bool,
96
97    #[serde(rename = "task.target_completion_interval")]
98    pub task_target_completion_interval: Option<String>,
99
100    #[serde(rename = "jdbc.url")]
101    pub jdbc_url: Option<String>,
102
103    #[serde(rename = "username")]
104    pub username: Option<String>,
105
106    #[serde(rename = "password")]
107    pub password: Option<String>,
108
109    // Authentication method control (password | key_pair_file | key_pair_object)
110    #[serde(rename = "auth.method")]
111    pub auth_method: Option<String>,
112
113    // Key-pair authentication via connection Properties (Option 2: file-based)
114    #[serde(rename = "private_key_file")]
115    pub private_key_file: Option<String>,
116
117    #[serde(rename = "private_key_file_pwd")]
118    pub private_key_file_pwd: Option<String>,
119
120    // Key-pair authentication via connection Properties (Option 1: object-based, PEM content)
121    #[serde(rename = "private_key_pem")]
122    pub private_key_pem: Option<String>,
123
124    /// Commit every n(>0) checkpoints, default is 10.
125    #[serde(default = "default_commit_checkpoint_interval")]
126    #[serde_as(as = "DisplayFromStr")]
127    #[with_option(allow_alter_on_fly)]
128    pub commit_checkpoint_interval: u64,
129
130    /// Enable auto schema change for upsert sink.
131    /// If enabled, the sink will automatically alter the target table to add new columns.
132    #[serde(default)]
133    #[serde(rename = "auto.schema.change")]
134    #[serde_as(as = "DisplayFromStr")]
135    pub auto_schema_change: bool,
136
137    #[serde(default)]
138    #[serde(rename = "create_table_if_not_exists")]
139    #[serde_as(as = "DisplayFromStr")]
140    pub create_table_if_not_exists: bool,
141
142    #[serde(default = "default_with_s3")]
143    #[serde(rename = "with_s3")]
144    #[serde_as(as = "DisplayFromStr")]
145    pub with_s3: bool,
146
147    #[serde(flatten)]
148    pub s3_inner: Option<S3Common>,
149
150    #[serde(rename = "stage")]
151    pub stage: Option<String>,
152
153    #[serde(flatten)]
154    pub unknown_fields: std::collections::HashMap<String, String>,
155}
156
157crate::impl_sink_unknown_fields!(SnowflakeV2Config);
158
159fn default_target_interval_schedule() -> u64 {
160    3600 // Default to 1 hour
161}
162
163fn default_intermediate_interval_schedule() -> u64 {
164    1800 // Default to 0.5 hour
165}
166
167fn default_with_s3() -> bool {
168    true
169}
170
171impl SnowflakeV2Config {
172    /// Build JDBC Properties for the Snowflake JDBC connection (no URL parameters).
173    /// Returns (`jdbc_url`, `driver_properties`).
174    /// - `driver_properties` are transformed/used by the Java runner and passed to `DriverManager::getConnection(url, props)`
175    ///
176    /// Note: This method assumes the config has been validated by `from_btreemap`.
177    pub fn build_jdbc_connection_properties(&self) -> Result<(String, Vec<(String, String)>)> {
178        let jdbc_url = self
179            .jdbc_url
180            .clone()
181            .ok_or(SinkError::Config(anyhow!("jdbc.url is required")))?;
182        let username = self
183            .username
184            .clone()
185            .ok_or(SinkError::Config(anyhow!("username is required")))?;
186
187        let mut connection_properties: Vec<(String, String)> = vec![("user".to_owned(), username)];
188
189        // auth_method is guaranteed to be Some after validation in from_btreemap
190        match self.auth_method.as_deref().unwrap() {
191            AUTH_METHOD_PASSWORD => {
192                // password is guaranteed to exist by from_btreemap validation
193                connection_properties.push(("password".to_owned(), self.password.clone().unwrap()));
194            }
195            AUTH_METHOD_KEY_PAIR_FILE => {
196                // private_key_file is guaranteed to exist by from_btreemap validation
197                connection_properties.push((
198                    "private_key_file".to_owned(),
199                    self.private_key_file.clone().unwrap(),
200                ));
201                if let Some(pwd) = self.private_key_file_pwd.clone() {
202                    connection_properties.push(("private_key_file_pwd".to_owned(), pwd));
203                }
204            }
205            AUTH_METHOD_KEY_PAIR_OBJECT => {
206                connection_properties.push((
207                    PROP_AUTH_METHOD.to_owned(),
208                    AUTH_METHOD_KEY_PAIR_OBJECT.to_owned(),
209                ));
210                // private_key_pem is guaranteed to exist by from_btreemap validation
211                connection_properties.push((
212                    "private_key_pem".to_owned(),
213                    self.private_key_pem.clone().unwrap(),
214                ));
215                if let Some(pwd) = self.private_key_file_pwd.clone() {
216                    connection_properties.push(("private_key_file_pwd".to_owned(), pwd));
217                }
218            }
219            _ => {
220                // This should never happen since from_btreemap validates auth_method
221                unreachable!(
222                    "Invalid auth_method - should have been caught during config validation"
223                )
224            }
225        }
226
227        Ok((jdbc_url, connection_properties))
228    }
229
230    pub fn from_btreemap(properties: &BTreeMap<String, String>) -> Result<Self> {
231        let mut config =
232            serde_json::from_value::<SnowflakeV2Config>(serde_json::to_value(properties).unwrap())
233                .map_err(|e| SinkError::Config(anyhow!(e)))?;
234        if config.r#type != SINK_TYPE_APPEND_ONLY && config.r#type != SINK_TYPE_UPSERT {
235            return Err(SinkError::Config(anyhow!(
236                "`{}` must be {}, or {}",
237                SINK_TYPE_OPTION,
238                SINK_TYPE_APPEND_ONLY,
239                SINK_TYPE_UPSERT
240            )));
241        }
242        let has_upsert_task_config = config.snowflake_cdc_table_name.is_some()
243            || properties.contains_key("write.target.interval.seconds")
244            || config.snowflake_warehouse.is_some()
245            || config.task_serverless
246            || config.task_target_completion_interval.is_some();
247        if config.r#type != SINK_TYPE_UPSERT && has_upsert_task_config {
248            return Err(SinkError::Config(anyhow!(
249                "`intermediate.table.name`, `write.target.interval.seconds`, `warehouse`, \
250                 `task.serverless`, and `task.target_completion_interval` require `{}` = {}",
251                SINK_TYPE_OPTION,
252                SINK_TYPE_UPSERT
253            )));
254        }
255        if config.task_target_completion_interval.is_some() && !config.task_serverless {
256            return Err(SinkError::Config(anyhow!(
257                "`task.target_completion_interval` requires `task.serverless` to be true"
258            )));
259        }
260        if config.task_serverless && config.snowflake_warehouse.is_some() {
261            return Err(SinkError::Config(anyhow!(
262                "`task.serverless` must not be combined with `warehouse`"
263            )));
264        }
265
266        // Normalize and validate authentication method
267        let has_password = config.password.is_some();
268        let has_file = config.private_key_file.is_some();
269        let has_pem = config.private_key_pem.as_deref().is_some();
270
271        let normalized_auth_method = match config
272            .auth_method
273            .as_deref()
274            .map(|s| s.trim().to_ascii_lowercase())
275        {
276            Some(method) if method == AUTH_METHOD_PASSWORD => {
277                if !has_password {
278                    return Err(SinkError::Config(anyhow!(
279                        "auth.method=password requires `password`"
280                    )));
281                }
282                if has_file || has_pem {
283                    return Err(SinkError::Config(anyhow!(
284                        "auth.method=password must not set `private_key_file`/`private_key_pem`"
285                    )));
286                }
287                AUTH_METHOD_PASSWORD.to_owned()
288            }
289            Some(method) if method == AUTH_METHOD_KEY_PAIR_FILE => {
290                if !has_file {
291                    return Err(SinkError::Config(anyhow!(
292                        "auth.method=key_pair_file requires `private_key_file`"
293                    )));
294                }
295                if has_password {
296                    return Err(SinkError::Config(anyhow!(
297                        "auth.method=key_pair_file must not set `password`"
298                    )));
299                }
300                if has_pem {
301                    return Err(SinkError::Config(anyhow!(
302                        "auth.method=key_pair_file must not set `private_key_pem`"
303                    )));
304                }
305                AUTH_METHOD_KEY_PAIR_FILE.to_owned()
306            }
307            Some(method) if method == AUTH_METHOD_KEY_PAIR_OBJECT => {
308                if !has_pem {
309                    return Err(SinkError::Config(anyhow!(
310                        "auth.method=key_pair_object requires `private_key_pem`"
311                    )));
312                }
313                if has_password {
314                    return Err(SinkError::Config(anyhow!(
315                        "auth.method=key_pair_object must not set `password`"
316                    )));
317                }
318                AUTH_METHOD_KEY_PAIR_OBJECT.to_owned()
319            }
320            Some(other) => {
321                return Err(SinkError::Config(anyhow!(
322                    "invalid auth.method: {} (allowed: password | key_pair_file | key_pair_object)",
323                    other
324                )));
325            }
326            None => {
327                // Infer auth method from supplied fields
328                match (has_password, has_file, has_pem) {
329                    (true, false, false) => AUTH_METHOD_PASSWORD.to_owned(),
330                    (false, true, false) => AUTH_METHOD_KEY_PAIR_FILE.to_owned(),
331                    (false, false, true) => AUTH_METHOD_KEY_PAIR_OBJECT.to_owned(),
332                    (true, true, _) | (true, _, true) | (false, true, true) => {
333                        return Err(SinkError::Config(anyhow!(
334                            "ambiguous auth: multiple auth options provided; remove one or set `auth.method`"
335                        )));
336                    }
337                    _ => {
338                        return Err(SinkError::Config(anyhow!(
339                            "no authentication configured: set either `password`, or `private_key_file`, or `private_key_pem` (or provide `auth.method`)"
340                        )));
341                    }
342                }
343            }
344        };
345        config.auth_method = Some(normalized_auth_method);
346        Ok(config)
347    }
348
349    pub fn build_snowflake_task_ctx_jdbc_client(
350        &self,
351        is_append_only: bool,
352        schema: &Schema,
353        pk_indices: &Vec<usize>,
354    ) -> Result<Option<(SnowflakeTaskContext, JdbcJniClient)>> {
355        if !self.auto_schema_change
356            && is_append_only
357            && !self.create_table_if_not_exists
358            && !self.with_s3
359        {
360            // append-only + no auto schema change is not need to create a client
361            return Ok(None);
362        }
363        let target_table_name = self
364            .snowflake_target_table_name
365            .clone()
366            .ok_or(SinkError::Config(anyhow!("table.name is required")))?;
367        let database = self
368            .snowflake_database
369            .clone()
370            .ok_or(SinkError::Config(anyhow!("database is required")))?;
371        let schema_name = self
372            .snowflake_schema
373            .clone()
374            .ok_or(SinkError::Config(anyhow!("schema is required")))?;
375        let mut snowflake_task_ctx = SnowflakeTaskContext {
376            target_table_name: target_table_name.clone(),
377            database,
378            schema_name,
379            schema: schema.clone(),
380            ..Default::default()
381        };
382
383        let (jdbc_url, connection_properties) = self.build_jdbc_connection_properties()?;
384        let client = JdbcJniClient::new_with_props(jdbc_url, connection_properties)?;
385
386        if self.with_s3 {
387            let stage = self
388                .stage
389                .clone()
390                .ok_or(SinkError::Config(anyhow!("stage is required")))?;
391            snowflake_task_ctx.stage = Some(stage);
392            snowflake_task_ctx.pipe_name = Some(format!("{}_pipe", target_table_name));
393        }
394        if !is_append_only {
395            let cdc_table_name = self
396                .snowflake_cdc_table_name
397                .clone()
398                .ok_or(SinkError::Config(anyhow!(
399                    "intermediate.table.name is required"
400                )))?;
401            snowflake_task_ctx.cdc_table_name = Some(cdc_table_name.clone());
402            snowflake_task_ctx.writer_target_interval_seconds = self.writer_target_interval_seconds;
403            snowflake_task_ctx.task_serverless = self.task_serverless;
404            snowflake_task_ctx.task_target_completion_interval =
405                self.task_target_completion_interval.clone();
406            if !self.task_serverless {
407                snowflake_task_ctx.warehouse = Some(
408                    self.snowflake_warehouse
409                        .clone()
410                        .ok_or(SinkError::Config(anyhow!("warehouse is required")))?,
411                );
412            }
413            let pk_column_names: Vec<_> = schema
414                .fields
415                .iter()
416                .enumerate()
417                .filter(|(index, _)| pk_indices.contains(index))
418                .map(|(_, field)| field.name.clone())
419                .collect();
420            if pk_column_names.is_empty() {
421                return Err(SinkError::Config(anyhow!(
422                    "Primary key columns not found. Please set the `primary_key` column in the sink properties, or ensure that the sink contains the primary key columns from the upstream."
423                )));
424            }
425            snowflake_task_ctx.pk_column_names = Some(pk_column_names);
426            snowflake_task_ctx.all_column_names = Some(
427                schema
428                    .fields
429                    .iter()
430                    .map(|field| field.name.clone())
431                    .collect(),
432            );
433            snowflake_task_ctx.task_name = Some(format!(
434                "rw_snowflake_sink_from_{cdc_table_name}_to_{target_table_name}"
435            ));
436        }
437        Ok(Some((snowflake_task_ctx, client)))
438    }
439}
440
441impl EnforceSecret for SnowflakeV2Config {
442    const ENFORCE_SECRET_PROPERTIES: Set<&'static str> = phf_set! {
443        "username",
444        "password",
445        "jdbc.url",
446        // Key-pair authentication secrets
447        "private_key_file_pwd",
448        "private_key_pem",
449    };
450}
451
452#[derive(Clone, Debug)]
453pub struct SnowflakeV2Sink {
454    config: SnowflakeV2Config,
455    schema: Schema,
456    pk_indices: Vec<usize>,
457    is_append_only: bool,
458    param: SinkParam,
459}
460
461impl EnforceSecret for SnowflakeV2Sink {
462    fn enforce_secret<'a>(
463        prop_iter: impl Iterator<Item = &'a str>,
464    ) -> crate::sink::ConnectorResult<()> {
465        for prop in prop_iter {
466            SnowflakeV2Config::enforce_one(prop)?;
467        }
468        Ok(())
469    }
470}
471
472impl TryFrom<SinkParam> for SnowflakeV2Sink {
473    type Error = SinkError;
474
475    fn try_from(param: SinkParam) -> std::result::Result<Self, Self::Error> {
476        let schema = param.schema();
477        let config = SnowflakeV2Config::from_btreemap(&param.properties)?;
478        let is_append_only = param.sink_type.is_append_only();
479        let pk_indices = param.downstream_pk_or_empty();
480        Ok(Self {
481            config,
482            schema,
483            pk_indices,
484            is_append_only,
485            param,
486        })
487    }
488}
489
490impl Sink for SnowflakeV2Sink {
491    type LogSinker = CoordinatedLogSinker<SnowflakeSinkWriter>;
492
493    const SINK_NAME: &'static str = SNOWFLAKE_SINK_V2;
494
495    crate::impl_validate_sink_unknown_fields!();
496
497    async fn validate(&self) -> Result<()> {
498        risingwave_common::license::Feature::SnowflakeSink
499            .check_available()
500            .map_err(|e| anyhow::anyhow!(e))?;
501        if let Some((snowflake_task_ctx, client)) =
502            self.config.build_snowflake_task_ctx_jdbc_client(
503                self.is_append_only,
504                &self.schema,
505                &self.pk_indices,
506            )?
507        {
508            let client = SnowflakeJniClient::new(client, snowflake_task_ctx);
509            client.execute_create_table().await?;
510            client.execute_create_pipe().await?;
511        }
512
513        Ok(())
514    }
515
516    fn support_schema_change() -> bool {
517        true
518    }
519
520    fn validate_alter_config(config: &BTreeMap<String, String>) -> Result<()> {
521        SnowflakeV2Config::from_btreemap(config)?;
522        Ok(())
523    }
524
525    async fn new_log_sinker(
526        &self,
527        writer_param: crate::sink::SinkWriterParam,
528    ) -> Result<Self::LogSinker> {
529        let writer = SnowflakeSinkWriter::new(
530            self.config.clone(),
531            self.is_append_only,
532            writer_param.clone(),
533            self.param.clone(),
534        )
535        .await?;
536
537        let commit_checkpoint_interval =
538            NonZeroU64::new(self.config.commit_checkpoint_interval).expect(
539                "commit_checkpoint_interval should be greater than 0, and it should be checked in config validation",
540            );
541
542        CoordinatedLogSinker::new(
543            &writer_param,
544            self.param.clone(),
545            writer,
546            commit_checkpoint_interval,
547        )
548        .await
549    }
550
551    fn is_coordinated_sink(&self) -> bool {
552        true
553    }
554
555    async fn new_coordinator(
556        &self,
557        _iceberg_compact_stat_sender: Option<UnboundedSender<IcebergSinkCompactionUpdate>>,
558    ) -> Result<SinkCommitCoordinator> {
559        let coordinator = SnowflakeSinkCommitter::new(
560            self.config.clone(),
561            &self.schema,
562            &self.pk_indices,
563            self.is_append_only,
564            self.param.sink_id,
565        )?;
566        Ok(SinkCommitCoordinator::SinglePhase(Box::new(coordinator)))
567    }
568}
569
570pub enum SnowflakeSinkWriter {
571    S3(SnowflakeRedshiftSinkS3Writer),
572    Jdbc(SnowflakeRedshiftSinkJdbcWriter),
573}
574
575impl SnowflakeSinkWriter {
576    pub async fn new(
577        config: SnowflakeV2Config,
578        is_append_only: bool,
579        writer_param: SinkWriterParam,
580        param: SinkParam,
581    ) -> Result<Self> {
582        let schema = param.schema();
583        let database = config.snowflake_database.ok_or_else(|| {
584            SinkError::Config(anyhow!("database is required for Snowflake JDBC sink"))
585        })?;
586        let schema_name = config.snowflake_schema.ok_or_else(|| {
587            SinkError::Config(anyhow!("schema is required for Snowflake JDBC sink"))
588        })?;
589        let table_name = config.snowflake_target_table_name.ok_or_else(|| {
590            SinkError::Config(anyhow!("table.name is required for Snowflake JDBC sink"))
591        })?;
592        if config.with_s3 {
593            let s3_writer = SnowflakeRedshiftSinkS3Writer::new(
594                config.s3_inner.ok_or_else(|| {
595                    SinkError::Config(anyhow!(
596                        "S3 configuration is required for Snowflake S3 sink"
597                    ))
598                })?,
599                schema,
600                is_append_only,
601                table_name,
602            )?;
603            Ok(Self::S3(s3_writer))
604        } else {
605            let jdbc_writer = SnowflakeRedshiftSinkJdbcWriter::new(
606                is_append_only,
607                writer_param,
608                param,
609                build_full_table_name(&database, &schema_name, &table_name),
610            )
611            .await?;
612            Ok(Self::Jdbc(jdbc_writer))
613        }
614    }
615}
616
617#[async_trait]
618impl SinkWriter for SnowflakeSinkWriter {
619    type CommitMetadata = Option<SinkMetadata>;
620
621    async fn begin_epoch(&mut self, epoch: u64) -> Result<()> {
622        match self {
623            Self::S3(writer) => writer.begin_epoch(epoch),
624            Self::Jdbc(writer) => writer.begin_epoch(epoch).await,
625        }
626    }
627
628    async fn write_batch(&mut self, chunk: StreamChunk) -> Result<()> {
629        match self {
630            Self::S3(writer) => writer.write_batch(chunk).await,
631            Self::Jdbc(writer) => writer.write_batch(chunk).await,
632        }
633    }
634
635    async fn barrier(&mut self, is_checkpoint: bool) -> Result<Option<SinkMetadata>> {
636        match self {
637            Self::S3(writer) => {
638                writer.barrier(is_checkpoint).await?;
639            }
640            Self::Jdbc(writer) => {
641                writer.barrier(is_checkpoint).await?;
642            }
643        }
644        Ok(Some(SinkMetadata {
645            metadata: Some(sink_metadata::Metadata::Serialized(
646                risingwave_pb::connector_service::sink_metadata::SerializedMetadata {
647                    metadata: vec![],
648                },
649            )),
650        }))
651    }
652
653    async fn abort(&mut self) -> Result<()> {
654        if let Self::Jdbc(writer) = self {
655            writer.abort().await
656        } else {
657            Ok(())
658        }
659    }
660}
661
662#[derive(Default, Clone)]
663pub struct SnowflakeTaskContext {
664    // required for task creation
665    pub target_table_name: String,
666    pub database: String,
667    pub schema_name: String,
668    pub schema: Schema,
669
670    // only upsert
671    pub task_name: Option<String>,
672    pub cdc_table_name: Option<String>,
673    pub writer_target_interval_seconds: u64,
674    pub warehouse: Option<String>,
675    pub task_serverless: bool,
676    pub task_target_completion_interval: Option<String>,
677    pub pk_column_names: Option<Vec<String>>,
678    pub all_column_names: Option<Vec<String>>,
679
680    // only s3 writer
681    pub stage: Option<String>,
682    pub pipe_name: Option<String>,
683}
684pub struct SnowflakeSinkCommitter {
685    client: Option<SnowflakeJniClient>,
686    _periodic_task_handle: Option<tokio::task::JoinHandle<()>>,
687    shutdown_sender: Option<tokio::sync::mpsc::UnboundedSender<()>>,
688}
689
690impl SnowflakeSinkCommitter {
691    pub fn new(
692        config: SnowflakeV2Config,
693        schema: &Schema,
694        pk_indices: &Vec<usize>,
695        is_append_only: bool,
696        sink_id: SinkId,
697    ) -> Result<Self> {
698        let (client, periodic_task_handle, shutdown_sender) =
699            if let Some((snowflake_task_ctx, client)) =
700                config.build_snowflake_task_ctx_jdbc_client(is_append_only, schema, pk_indices)?
701            {
702                let (shutdown_sender, shutdown_receiver) = unbounded_channel();
703                let snowflake_client =
704                    SnowflakeJniClient::new(client.clone(), snowflake_task_ctx.clone());
705                let periodic_task_handle = tokio::spawn(async move {
706                    Self::run_periodic_query_task(
707                        snowflake_client,
708                        config.write_intermediate_interval_seconds,
709                        sink_id,
710                        shutdown_receiver,
711                    )
712                    .await;
713                });
714                (
715                    Some(SnowflakeJniClient::new(client, snowflake_task_ctx)),
716                    Some(periodic_task_handle),
717                    Some(shutdown_sender),
718                )
719            } else {
720                (None, None, None)
721            };
722
723        Ok(Self {
724            client,
725            _periodic_task_handle: periodic_task_handle,
726            shutdown_sender,
727        })
728    }
729
730    async fn run_periodic_query_task(
731        client: SnowflakeJniClient,
732        write_intermediate_interval_seconds: u64,
733        sink_id: SinkId,
734        mut shutdown_receiver: tokio::sync::mpsc::UnboundedReceiver<()>,
735    ) {
736        let mut copy_timer = interval(Duration::from_secs(write_intermediate_interval_seconds));
737        copy_timer.set_missed_tick_behavior(MissedTickBehavior::Skip);
738        loop {
739            tokio::select! {
740                _ = shutdown_receiver.recv() => break,
741                _ = copy_timer.tick() => {
742                    if let Err(e) = async {
743                        client.execute_flush_pipe().await?;
744                        Ok::<(),SinkError>(())
745                    }.await {
746                        tracing::error!("Failed to execute copy into task for sink id {}: {}", sink_id, e.as_report());
747                    }
748                }
749            }
750        }
751        tracing::info!("Periodic query task stopped for sink id {}", sink_id);
752    }
753}
754
755#[async_trait]
756impl SinglePhaseCommitCoordinator for SnowflakeSinkCommitter {
757    async fn init(&mut self) -> Result<()> {
758        if let Some(client) = &self.client {
759            // Todo: move this to validate
760            client.execute_create_pipe().await?;
761            client.execute_create_merge_into_task().await?;
762        }
763        Ok(())
764    }
765
766    async fn commit_data(&mut self, _epoch: u64, _metadata: Vec<SinkMetadata>) -> Result<()> {
767        Ok(())
768    }
769
770    async fn commit_schema_change(
771        &mut self,
772        _epoch: u64,
773        schema_change: PbSinkSchemaChange,
774    ) -> Result<()> {
775        use risingwave_pb::stream_plan::sink_schema_change::PbOp as SinkSchemaChangeOp;
776        let schema_change_op = schema_change
777            .op
778            .ok_or_else(|| SinkError::Coordinator(anyhow!("Invalid schema change operation")))?;
779        let SinkSchemaChangeOp::AddColumns(add_columns) = schema_change_op else {
780            return Err(SinkError::Coordinator(anyhow!(
781                "Only AddColumns schema change is supported for Snowflake sink"
782            )));
783        };
784        let client = self.client.as_mut().ok_or_else(|| {
785            SinkError::Config(anyhow!("Snowflake sink committer is not initialized."))
786        })?;
787        client
788            .execute_alter_add_columns(
789                &add_columns
790                    .fields
791                    .into_iter()
792                    .map(|f| {
793                        let dt = DataType::from(f.data_type.unwrap());
794                        Ok((f.name, convert_snowflake_data_type(&dt)?))
795                    })
796                    .collect::<Result<Vec<_>>>()?,
797            )
798            .await
799    }
800}
801
802impl Drop for SnowflakeSinkCommitter {
803    fn drop(&mut self) {
804        if let Some(client) = self.client.take() {
805            if let Some(sender) = self.shutdown_sender.take() {
806                let _ = sender.send(()); // Ignore the result, as the receiver may have been dropped.
807            }
808            tokio::spawn(async move {
809                client.execute_drop_task().await.ok();
810            });
811        }
812    }
813}
814
815pub struct SnowflakeJniClient {
816    jdbc_client: JdbcJniClient,
817    snowflake_task_context: SnowflakeTaskContext,
818}
819
820impl SnowflakeJniClient {
821    pub fn new(jdbc_client: JdbcJniClient, snowflake_task_context: SnowflakeTaskContext) -> Self {
822        Self {
823            jdbc_client,
824            snowflake_task_context,
825        }
826    }
827
828    pub async fn execute_alter_add_columns(
829        &mut self,
830        columns: &Vec<(String, String)>,
831    ) -> Result<()> {
832        self.execute_drop_task().await?;
833        if let Some(names) = self.snowflake_task_context.all_column_names.as_mut() {
834            names.extend(columns.iter().map(|(name, _)| name.clone()));
835        }
836        if let Some(cdc_table_name) = &self.snowflake_task_context.cdc_table_name {
837            let alter_add_column_cdc_table_sql = build_alter_add_column_sql(
838                cdc_table_name,
839                &self.snowflake_task_context.database,
840                &self.snowflake_task_context.schema_name,
841                columns,
842            );
843            self.jdbc_client
844                .execute_sql_sync(vec![alter_add_column_cdc_table_sql])
845                .await?;
846        }
847
848        let alter_add_column_target_table_sql = build_alter_add_column_sql(
849            &self.snowflake_task_context.target_table_name,
850            &self.snowflake_task_context.database,
851            &self.snowflake_task_context.schema_name,
852            columns,
853        );
854        self.jdbc_client
855            .execute_sql_sync(vec![alter_add_column_target_table_sql])
856            .await?;
857
858        self.execute_create_merge_into_task().await?;
859        Ok(())
860    }
861
862    pub async fn execute_create_merge_into_task(&self) -> Result<()> {
863        if self.snowflake_task_context.task_name.is_some() {
864            let create_task_sql = build_create_merge_into_task_sql(&self.snowflake_task_context);
865            let start_task_sql = build_start_task_sql(&self.snowflake_task_context);
866            self.jdbc_client
867                .execute_sql_sync(vec![create_task_sql])
868                .await?;
869            self.jdbc_client
870                .execute_sql_sync(vec![start_task_sql])
871                .await?;
872        }
873        Ok(())
874    }
875
876    pub async fn execute_drop_task(&self) -> Result<()> {
877        if self.snowflake_task_context.task_name.is_some() {
878            let sql = build_drop_task_sql(&self.snowflake_task_context);
879            if let Err(e) = self.jdbc_client.execute_sql_sync(vec![sql]).await {
880                tracing::error!(
881                    "Failed to drop Snowflake sink task {:?}: {:?}",
882                    self.snowflake_task_context.task_name,
883                    e.as_report()
884                );
885            } else {
886                tracing::info!(
887                    "Snowflake sink task {:?} dropped",
888                    self.snowflake_task_context.task_name
889                );
890            }
891        }
892        Ok(())
893    }
894
895    pub async fn execute_create_table(&self) -> Result<()> {
896        // create target table
897        let create_target_table_sql = build_create_table_sql(
898            &self.snowflake_task_context.target_table_name,
899            &self.snowflake_task_context.database,
900            &self.snowflake_task_context.schema_name,
901            &self.snowflake_task_context.schema,
902            false,
903        )?;
904        self.jdbc_client
905            .execute_sql_sync(vec![create_target_table_sql])
906            .await?;
907        if let Some(cdc_table_name) = &self.snowflake_task_context.cdc_table_name {
908            let create_cdc_table_sql = build_create_table_sql(
909                cdc_table_name,
910                &self.snowflake_task_context.database,
911                &self.snowflake_task_context.schema_name,
912                &self.snowflake_task_context.schema,
913                true,
914            )?;
915            self.jdbc_client
916                .execute_sql_sync(vec![create_cdc_table_sql])
917                .await?;
918        }
919        Ok(())
920    }
921
922    pub async fn execute_create_pipe(&self) -> Result<()> {
923        if let Some(pipe_name) = &self.snowflake_task_context.pipe_name {
924            let table_name =
925                if let Some(table_name) = self.snowflake_task_context.cdc_table_name.as_ref() {
926                    table_name
927                } else {
928                    &self.snowflake_task_context.target_table_name
929                };
930            let create_pipe_sql = build_create_pipe_sql(
931                table_name,
932                &self.snowflake_task_context.database,
933                &self.snowflake_task_context.schema_name,
934                self.snowflake_task_context.stage.as_ref().ok_or_else(|| {
935                    SinkError::Config(anyhow!("snowflake.stage is required for S3 writer"))
936                })?,
937                pipe_name,
938                &self.snowflake_task_context.target_table_name,
939            );
940            self.jdbc_client
941                .execute_sql_sync(vec![create_pipe_sql])
942                .await?;
943        }
944        Ok(())
945    }
946
947    pub async fn execute_flush_pipe(&self) -> Result<()> {
948        if let Some(pipe_name) = &self.snowflake_task_context.pipe_name {
949            let flush_pipe_sql = build_flush_pipe_sql(
950                &self.snowflake_task_context.database,
951                &self.snowflake_task_context.schema_name,
952                pipe_name,
953            );
954            self.jdbc_client
955                .execute_sql_sync(vec![flush_pipe_sql])
956                .await?;
957        }
958        Ok(())
959    }
960}
961
962fn build_create_table_sql(
963    table_name: &str,
964    database: &str,
965    schema_name: &str,
966    schema: &Schema,
967    need_op_and_row_id: bool,
968) -> Result<String> {
969    let full_table_name = build_full_table_name(database, schema_name, table_name);
970    let mut columns: Vec<String> = schema
971        .fields
972        .iter()
973        .map(|field| {
974            let data_type = convert_snowflake_data_type(&field.data_type)?;
975            Ok(format!(r#""{}" {}"#, field.name, data_type))
976        })
977        .collect::<Result<Vec<String>>>()?;
978    if need_op_and_row_id {
979        columns.push(format!(r#""{}" STRING"#, __ROW_ID));
980        columns.push(format!(r#""{}" INT"#, __OP));
981    }
982    let columns_str = columns.join(", ");
983    Ok(format!(
984        "CREATE TABLE IF NOT EXISTS {} ({}) ENABLE_SCHEMA_EVOLUTION  = true",
985        full_table_name, columns_str
986    ))
987}
988
989fn convert_snowflake_data_type(data_type: &DataType) -> Result<String> {
990    let data_type = match data_type {
991        DataType::Int16 => "SMALLINT".to_owned(),
992        DataType::Int32 => "INTEGER".to_owned(),
993        DataType::Int64 => "BIGINT".to_owned(),
994        DataType::Float32 => "FLOAT4".to_owned(),
995        DataType::Float64 => "FLOAT8".to_owned(),
996        DataType::Boolean => "BOOLEAN".to_owned(),
997        DataType::Varchar => "STRING".to_owned(),
998        DataType::Date => "DATE".to_owned(),
999        DataType::Timestamp => "TIMESTAMP".to_owned(),
1000        DataType::Timestamptz => "TIMESTAMP_TZ".to_owned(),
1001        DataType::Jsonb => "STRING".to_owned(),
1002        // RisingWave uses rust_decimal with MAX_PRECISION=28. Snowflake's DECIMAL without
1003        // explicit precision defaults to (38,0), which drops all fractional digits. We use
1004        // DECIMAL(38, 10) to preserve up to 10 fractional digits, matching the Iceberg sink
1005        // convention, though values with more than 10 fractional digits may still lose precision.
1006        DataType::Decimal => "DECIMAL(38, 10)".to_owned(),
1007        DataType::Bytea => "BINARY".to_owned(),
1008        DataType::Time => "TIME".to_owned(),
1009        _ => {
1010            return Err(SinkError::Config(anyhow!(
1011                "Dont support auto create table for datatype: {}",
1012                data_type
1013            )));
1014        }
1015    };
1016    Ok(data_type)
1017}
1018
1019fn build_create_pipe_sql(
1020    table_name: &str,
1021    database: &str,
1022    schema: &str,
1023    stage: &str,
1024    pipe_name: &str,
1025    target_table_name: &str,
1026) -> String {
1027    let pipe_name = format!(r#""{}"."{}"."{}""#, database, schema, pipe_name);
1028    // Trailing `/` is required to enforce exact directory matching.
1029    // Without it, a PIPE for table "dim_project" would also match files under
1030    // "dim_project_contract/" since Snowflake uses prefix matching.
1031    let stage = format!(
1032        r#""{}"."{}"."{}"/{}/"#,
1033        database, schema, stage, target_table_name
1034    );
1035    let table_name = format!(r#""{}"."{}"."{}""#, database, schema, table_name);
1036    format!(
1037        "CREATE OR REPLACE PIPE {} AUTO_INGEST = FALSE AS COPY INTO {} FROM @{} MATCH_BY_COLUMN_NAME = CASE_INSENSITIVE FILE_FORMAT = (type = 'JSON');",
1038        pipe_name, table_name, stage
1039    )
1040}
1041
1042fn build_flush_pipe_sql(database: &str, schema: &str, pipe_name: &str) -> String {
1043    let pipe_name = format!(r#""{}"."{}"."{}""#, database, schema, pipe_name);
1044    format!("ALTER PIPE {} REFRESH;", pipe_name,)
1045}
1046
1047fn build_alter_add_column_sql(
1048    table_name: &str,
1049    database: &str,
1050    schema: &str,
1051    columns: &Vec<(String, String)>,
1052) -> String {
1053    let full_table_name = build_full_table_name(database, schema, table_name);
1054    jdbc_jni_client::build_alter_add_column_sql(&full_table_name, columns, true)
1055}
1056
1057fn build_start_task_sql(snowflake_task_context: &SnowflakeTaskContext) -> String {
1058    let SnowflakeTaskContext {
1059        task_name,
1060        database,
1061        schema_name: schema,
1062        ..
1063    } = snowflake_task_context;
1064    let full_task_name = format!(
1065        r#""{}"."{}"."{}""#,
1066        database,
1067        schema,
1068        task_name.as_ref().unwrap()
1069    );
1070    format!("ALTER TASK {} RESUME", full_task_name)
1071}
1072
1073fn build_drop_task_sql(snowflake_task_context: &SnowflakeTaskContext) -> String {
1074    let SnowflakeTaskContext {
1075        task_name,
1076        database,
1077        schema_name: schema,
1078        ..
1079    } = snowflake_task_context;
1080    let full_task_name = format!(
1081        r#""{}"."{}"."{}""#,
1082        database,
1083        schema,
1084        task_name.as_ref().unwrap()
1085    );
1086    format!("DROP TASK IF EXISTS {}", full_task_name)
1087}
1088
1089fn build_create_merge_into_task_sql(snowflake_task_context: &SnowflakeTaskContext) -> String {
1090    let SnowflakeTaskContext {
1091        task_name,
1092        cdc_table_name,
1093        target_table_name,
1094        writer_target_interval_seconds,
1095        warehouse,
1096        task_serverless,
1097        task_target_completion_interval,
1098        pk_column_names,
1099        all_column_names,
1100        database,
1101        schema_name,
1102        ..
1103    } = snowflake_task_context;
1104    let full_task_name = format!(
1105        r#""{}"."{}"."{}""#,
1106        database,
1107        schema_name,
1108        task_name.as_ref().unwrap()
1109    );
1110    let full_cdc_table_name = format!(
1111        r#""{}"."{}"."{}""#,
1112        database,
1113        schema_name,
1114        cdc_table_name.as_ref().unwrap()
1115    );
1116    let full_target_table_name = format!(
1117        r#""{}"."{}"."{}""#,
1118        database, schema_name, target_table_name
1119    );
1120
1121    let pk_names_str = pk_column_names
1122        .as_ref()
1123        .unwrap()
1124        .iter()
1125        .map(|name| format!(r#""{}""#, name))
1126        .collect::<Vec<String>>()
1127        .join(", ");
1128    let pk_names_eq_str = pk_column_names
1129        .as_ref()
1130        .unwrap()
1131        .iter()
1132        .map(|name| format!(r#"target."{}" = source."{}""#, name, name))
1133        .collect::<Vec<String>>()
1134        .join(" AND ");
1135    let all_column_names_set_str = all_column_names
1136        .as_ref()
1137        .unwrap()
1138        .iter()
1139        .map(|name| format!(r#"target."{}" = source."{}""#, name, name))
1140        .collect::<Vec<String>>()
1141        .join(", ");
1142    let all_column_names_str = all_column_names
1143        .as_ref()
1144        .unwrap()
1145        .iter()
1146        .map(|name| format!(r#""{}""#, name))
1147        .collect::<Vec<String>>()
1148        .join(", ");
1149    let all_column_names_insert_str = all_column_names
1150        .as_ref()
1151        .unwrap()
1152        .iter()
1153        .map(|name| format!(r#"source."{}""#, name))
1154        .collect::<Vec<String>>()
1155        .join(", ");
1156
1157    let compute_clause = if *task_serverless {
1158        task_target_completion_interval
1159            .as_ref()
1160            .map(|interval| format!("TARGET_COMPLETION_INTERVAL = '{}'", interval))
1161    } else {
1162        Some(format!("WAREHOUSE = {}", warehouse.as_ref().unwrap()))
1163    };
1164
1165    format!(
1166        r#"CREATE OR REPLACE TASK {task_name}
1167{compute_clause}
1168SCHEDULE = '{writer_target_interval_seconds} SECONDS'
1169AS
1170BEGIN
1171    LET max_row_id STRING;
1172
1173    SELECT COALESCE(MAX("{snowflake_sink_row_id}"), '0') INTO :max_row_id
1174    FROM {cdc_table_name};
1175
1176    MERGE INTO {target_table_name} AS target
1177    USING (
1178        SELECT *
1179        FROM (
1180            SELECT *, ROW_NUMBER() OVER (PARTITION BY {pk_names_str} ORDER BY "{snowflake_sink_row_id}" DESC) AS dedupe_id
1181            FROM {cdc_table_name}
1182            WHERE "{snowflake_sink_row_id}" <= :max_row_id
1183        ) AS subquery
1184        WHERE dedupe_id = 1
1185    ) AS source
1186    ON {pk_names_eq_str}
1187    WHEN MATCHED AND source."{snowflake_sink_op}" IN (2, 4) THEN DELETE
1188    WHEN MATCHED AND source."{snowflake_sink_op}" IN (1, 3) THEN UPDATE SET {all_column_names_set_str}
1189    WHEN NOT MATCHED AND source."{snowflake_sink_op}" IN (1, 3) THEN INSERT ({all_column_names_str}) VALUES ({all_column_names_insert_str});
1190
1191    DELETE FROM {cdc_table_name}
1192    WHERE "{snowflake_sink_row_id}" <= :max_row_id;
1193END;"#,
1194        task_name = full_task_name,
1195        compute_clause = compute_clause
1196            .map(|clause| format!("{clause}\n"))
1197            .unwrap_or_default(),
1198        writer_target_interval_seconds = writer_target_interval_seconds,
1199        cdc_table_name = full_cdc_table_name,
1200        target_table_name = full_target_table_name,
1201        pk_names_str = pk_names_str,
1202        pk_names_eq_str = pk_names_eq_str,
1203        all_column_names_set_str = all_column_names_set_str,
1204        all_column_names_str = all_column_names_str,
1205        all_column_names_insert_str = all_column_names_insert_str,
1206        snowflake_sink_row_id = __ROW_ID,
1207        snowflake_sink_op = __OP,
1208    )
1209}
1210
1211#[cfg(test)]
1212mod tests {
1213    use std::collections::BTreeMap;
1214
1215    use super::*;
1216    use crate::sink::jdbc_jni_client::normalize_sql;
1217
1218    fn base_properties() -> BTreeMap<String, String> {
1219        BTreeMap::from([
1220            ("type".to_owned(), "append-only".to_owned()),
1221            ("jdbc.url".to_owned(), "jdbc:snowflake://account".to_owned()),
1222            ("username".to_owned(), "RW_USER".to_owned()),
1223        ])
1224    }
1225
1226    #[test]
1227    fn test_build_jdbc_props_password() {
1228        let mut props = base_properties();
1229        props.insert("password".to_owned(), "secret".to_owned());
1230        let config = SnowflakeV2Config::from_btreemap(&props).unwrap();
1231        let (url, connection_properties) = config.build_jdbc_connection_properties().unwrap();
1232        assert_eq!(url, "jdbc:snowflake://account");
1233        let map: BTreeMap<_, _> = connection_properties.into_iter().collect();
1234        assert_eq!(map.get("user"), Some(&"RW_USER".to_owned()));
1235        assert_eq!(map.get("password"), Some(&"secret".to_owned()));
1236        assert!(!map.contains_key("authenticator"));
1237    }
1238
1239    #[test]
1240    fn test_build_jdbc_props_key_pair_file() {
1241        let mut props = base_properties();
1242        props.insert(
1243            "auth.method".to_owned(),
1244            AUTH_METHOD_KEY_PAIR_FILE.to_owned(),
1245        );
1246        props.insert("private_key_file".to_owned(), "/tmp/rsa_key.p8".to_owned());
1247        props.insert("private_key_file_pwd".to_owned(), "dummy".to_owned());
1248        let config = SnowflakeV2Config::from_btreemap(&props).unwrap();
1249        let (url, connection_properties) = config.build_jdbc_connection_properties().unwrap();
1250        assert_eq!(url, "jdbc:snowflake://account");
1251        let map: BTreeMap<_, _> = connection_properties.into_iter().collect();
1252        assert_eq!(map.get("user"), Some(&"RW_USER".to_owned()));
1253        assert_eq!(
1254            map.get("private_key_file"),
1255            Some(&"/tmp/rsa_key.p8".to_owned())
1256        );
1257        assert_eq!(map.get("private_key_file_pwd"), Some(&"dummy".to_owned()));
1258    }
1259
1260    #[test]
1261    fn test_build_jdbc_props_key_pair_object() {
1262        let mut props = base_properties();
1263        props.insert(
1264            "auth.method".to_owned(),
1265            AUTH_METHOD_KEY_PAIR_OBJECT.to_owned(),
1266        );
1267        props.insert(
1268            "private_key_pem".to_owned(),
1269            "-----BEGIN PRIVATE KEY-----
1270...
1271-----END PRIVATE KEY-----"
1272                .to_owned(),
1273        );
1274        let config = SnowflakeV2Config::from_btreemap(&props).unwrap();
1275        let (url, connection_properties) = config.build_jdbc_connection_properties().unwrap();
1276        assert_eq!(url, "jdbc:snowflake://account");
1277        let map: BTreeMap<_, _> = connection_properties.into_iter().collect();
1278        assert_eq!(
1279            map.get("private_key_pem"),
1280            Some(
1281                &"-----BEGIN PRIVATE KEY-----
1282...
1283-----END PRIVATE KEY-----"
1284                    .to_owned()
1285            )
1286        );
1287        assert!(!map.contains_key("private_key_file"));
1288    }
1289
1290    #[test]
1291    fn test_snowflake_task_target_completion_interval_requires_serverless() {
1292        let mut props = base_properties();
1293        props.insert("password".to_owned(), "secret".to_owned());
1294        props.insert("type".to_owned(), "upsert".to_owned());
1295        props.insert(
1296            "task.target_completion_interval".to_owned(),
1297            "5 MINUTES".to_owned(),
1298        );
1299
1300        let err = SnowflakeV2Config::from_btreemap(&props).unwrap_err();
1301        assert!(
1302            err.as_report().to_string().contains(
1303                "`task.target_completion_interval` requires `task.serverless` to be true"
1304            )
1305        );
1306    }
1307
1308    #[test]
1309    fn test_snowflake_serverless_task_rejects_warehouse() {
1310        let mut props = base_properties();
1311        props.insert("password".to_owned(), "secret".to_owned());
1312        props.insert("type".to_owned(), "upsert".to_owned());
1313        props.insert("task.serverless".to_owned(), "true".to_owned());
1314        props.insert("warehouse".to_owned(), "test_warehouse".to_owned());
1315
1316        let err = SnowflakeV2Config::from_btreemap(&props).unwrap_err();
1317        assert!(
1318            err.as_report()
1319                .to_string()
1320                .contains("`task.serverless` must not be combined with `warehouse`")
1321        );
1322    }
1323
1324    #[test]
1325    fn test_snowflake_append_only_rejects_upsert_task_options() {
1326        for (key, value) in [
1327            ("intermediate.table.name", "test_intermediate"),
1328            ("write.target.interval.seconds", "3600"),
1329            ("warehouse", "test_warehouse"),
1330            ("task.serverless", "true"),
1331            ("task.target_completion_interval", "5 MINUTES"),
1332        ] {
1333            let mut props = base_properties();
1334            props.insert("password".to_owned(), "secret".to_owned());
1335            props.insert(key.to_owned(), value.to_owned());
1336
1337            let err = SnowflakeV2Config::from_btreemap(&props).unwrap_err();
1338            assert!(
1339                err.as_report().to_string().contains(
1340                    "`intermediate.table.name`, `write.target.interval.seconds`, `warehouse`, \
1341                 `task.serverless`, and `task.target_completion_interval` require `type` = upsert"
1342                ),
1343                "option {key} should be rejected for append-only sink"
1344            );
1345        }
1346    }
1347
1348    #[test]
1349    fn test_snowflake_sink_commit_coordinator() {
1350        let snowflake_task_context = SnowflakeTaskContext {
1351            task_name: Some("test_task".to_owned()),
1352            cdc_table_name: Some("test_cdc_table".to_owned()),
1353            target_table_name: "test_target_table".to_owned(),
1354            writer_target_interval_seconds: 3600,
1355            warehouse: Some("test_warehouse".to_owned()),
1356            task_serverless: false,
1357            task_target_completion_interval: None,
1358            pk_column_names: Some(vec!["v1".to_owned()]),
1359            all_column_names: Some(vec!["v1".to_owned(), "v2".to_owned()]),
1360            database: "test_db".to_owned(),
1361            schema_name: "test_schema".to_owned(),
1362            schema: Schema { fields: vec![] },
1363            stage: None,
1364            pipe_name: None,
1365        };
1366        let task_sql = build_create_merge_into_task_sql(&snowflake_task_context);
1367        let expected = r#"CREATE OR REPLACE TASK "test_db"."test_schema"."test_task"
1368WAREHOUSE = test_warehouse
1369SCHEDULE = '3600 SECONDS'
1370AS
1371BEGIN
1372    LET max_row_id STRING;
1373
1374    SELECT COALESCE(MAX("__row_id"), '0') INTO :max_row_id
1375    FROM "test_db"."test_schema"."test_cdc_table";
1376
1377    MERGE INTO "test_db"."test_schema"."test_target_table" AS target
1378    USING (
1379        SELECT *
1380        FROM (
1381            SELECT *, ROW_NUMBER() OVER (PARTITION BY "v1" ORDER BY "__row_id" DESC) AS dedupe_id
1382            FROM "test_db"."test_schema"."test_cdc_table"
1383            WHERE "__row_id" <= :max_row_id
1384        ) AS subquery
1385        WHERE dedupe_id = 1
1386    ) AS source
1387    ON target."v1" = source."v1"
1388    WHEN MATCHED AND source."__op" IN (2, 4) THEN DELETE
1389    WHEN MATCHED AND source."__op" IN (1, 3) THEN UPDATE SET target."v1" = source."v1", target."v2" = source."v2"
1390    WHEN NOT MATCHED AND source."__op" IN (1, 3) THEN INSERT ("v1", "v2") VALUES (source."v1", source."v2");
1391
1392    DELETE FROM "test_db"."test_schema"."test_cdc_table"
1393    WHERE "__row_id" <= :max_row_id;
1394END;"#;
1395        assert_eq!(normalize_sql(&task_sql), normalize_sql(expected));
1396    }
1397
1398    #[test]
1399    fn test_snowflake_sink_commit_coordinator_multi_pk() {
1400        let snowflake_task_context = SnowflakeTaskContext {
1401            task_name: Some("test_task_multi_pk".to_owned()),
1402            cdc_table_name: Some("cdc_multi_pk".to_owned()),
1403            target_table_name: "target_multi_pk".to_owned(),
1404            writer_target_interval_seconds: 300,
1405            warehouse: Some("multi_pk_warehouse".to_owned()),
1406            task_serverless: false,
1407            task_target_completion_interval: None,
1408            pk_column_names: Some(vec!["id1".to_owned(), "id2".to_owned()]),
1409            all_column_names: Some(vec!["id1".to_owned(), "id2".to_owned(), "val".to_owned()]),
1410            database: "test_db".to_owned(),
1411            schema_name: "test_schema".to_owned(),
1412            schema: Schema { fields: vec![] },
1413            stage: None,
1414            pipe_name: None,
1415        };
1416        let task_sql = build_create_merge_into_task_sql(&snowflake_task_context);
1417        let expected = r#"CREATE OR REPLACE TASK "test_db"."test_schema"."test_task_multi_pk"
1418WAREHOUSE = multi_pk_warehouse
1419SCHEDULE = '300 SECONDS'
1420AS
1421BEGIN
1422    LET max_row_id STRING;
1423
1424    SELECT COALESCE(MAX("__row_id"), '0') INTO :max_row_id
1425    FROM "test_db"."test_schema"."cdc_multi_pk";
1426
1427    MERGE INTO "test_db"."test_schema"."target_multi_pk" AS target
1428    USING (
1429        SELECT *
1430        FROM (
1431            SELECT *, ROW_NUMBER() OVER (PARTITION BY "id1", "id2" ORDER BY "__row_id" DESC) AS dedupe_id
1432            FROM "test_db"."test_schema"."cdc_multi_pk"
1433            WHERE "__row_id" <= :max_row_id
1434        ) AS subquery
1435        WHERE dedupe_id = 1
1436    ) AS source
1437    ON target."id1" = source."id1" AND target."id2" = source."id2"
1438    WHEN MATCHED AND source."__op" IN (2, 4) THEN DELETE
1439    WHEN MATCHED AND source."__op" IN (1, 3) THEN UPDATE SET target."id1" = source."id1", target."id2" = source."id2", target."val" = source."val"
1440    WHEN NOT MATCHED AND source."__op" IN (1, 3) THEN INSERT ("id1", "id2", "val") VALUES (source."id1", source."id2", source."val");
1441
1442    DELETE FROM "test_db"."test_schema"."cdc_multi_pk"
1443    WHERE "__row_id" <= :max_row_id;
1444END;"#;
1445        assert_eq!(normalize_sql(&task_sql), normalize_sql(expected));
1446    }
1447
1448    #[test]
1449    fn test_snowflake_sink_commit_coordinator_serverless_task() {
1450        let snowflake_task_context = SnowflakeTaskContext {
1451            task_name: Some("test_serverless_task".to_owned()),
1452            cdc_table_name: Some("serverless_cdc_table".to_owned()),
1453            target_table_name: "serverless_target_table".to_owned(),
1454            writer_target_interval_seconds: 120,
1455            warehouse: None,
1456            task_serverless: true,
1457            task_target_completion_interval: Some("5 MINUTES".to_owned()),
1458            pk_column_names: Some(vec!["id".to_owned()]),
1459            all_column_names: Some(vec!["id".to_owned(), "val".to_owned()]),
1460            database: "test_db".to_owned(),
1461            schema_name: "test_schema".to_owned(),
1462            schema: Schema { fields: vec![] },
1463            stage: None,
1464            pipe_name: None,
1465        };
1466        let task_sql = build_create_merge_into_task_sql(&snowflake_task_context);
1467        let expected = r#"CREATE OR REPLACE TASK "test_db"."test_schema"."test_serverless_task"
1468TARGET_COMPLETION_INTERVAL = '5 MINUTES'
1469SCHEDULE = '120 SECONDS'
1470AS
1471BEGIN
1472    LET max_row_id STRING;
1473
1474    SELECT COALESCE(MAX("__row_id"), '0') INTO :max_row_id
1475    FROM "test_db"."test_schema"."serverless_cdc_table";
1476
1477    MERGE INTO "test_db"."test_schema"."serverless_target_table" AS target
1478    USING (
1479        SELECT *
1480        FROM (
1481            SELECT *, ROW_NUMBER() OVER (PARTITION BY "id" ORDER BY "__row_id" DESC) AS dedupe_id
1482            FROM "test_db"."test_schema"."serverless_cdc_table"
1483            WHERE "__row_id" <= :max_row_id
1484        ) AS subquery
1485        WHERE dedupe_id = 1
1486    ) AS source
1487    ON target."id" = source."id"
1488    WHEN MATCHED AND source."__op" IN (2, 4) THEN DELETE
1489    WHEN MATCHED AND source."__op" IN (1, 3) THEN UPDATE SET target."id" = source."id", target."val" = source."val"
1490    WHEN NOT MATCHED AND source."__op" IN (1, 3) THEN INSERT ("id", "val") VALUES (source."id", source."val");
1491
1492    DELETE FROM "test_db"."test_schema"."serverless_cdc_table"
1493    WHERE "__row_id" <= :max_row_id;
1494END;"#;
1495        assert_eq!(normalize_sql(&task_sql), normalize_sql(expected));
1496    }
1497
1498    #[test]
1499    fn test_build_create_pipe_sql_stage_has_trailing_slash() {
1500        let sql = build_create_pipe_sql(
1501            "reservations_intermediate",
1502            "test_db",
1503            "test_schema",
1504            "RW_S3_STAGE",
1505            "reservations_pipe",
1506            "reservations",
1507        );
1508        assert!(
1509            sql.contains(r#"FROM @"test_db"."test_schema"."RW_S3_STAGE"/reservations/ "#),
1510            "unexpected pipe sql: {sql}"
1511        );
1512    }
1513
1514    #[test]
1515    fn test_convert_snowflake_decimal_data_type() {
1516        assert_eq!(
1517            convert_snowflake_data_type(&DataType::Decimal).unwrap(),
1518            "DECIMAL(38, 10)"
1519        );
1520    }
1521}