1use 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 #[serde(rename = "auth.method")]
111 pub auth_method: Option<String>,
112
113 #[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 #[serde(rename = "private_key_pem")]
122 pub private_key_pem: Option<String>,
123
124 #[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 #[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 }
162
163fn default_intermediate_interval_schedule() -> u64 {
164 1800 }
166
167fn default_with_s3() -> bool {
168 true
169}
170
171impl SnowflakeV2Config {
172 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 match self.auth_method.as_deref().unwrap() {
191 AUTH_METHOD_PASSWORD => {
192 connection_properties.push(("password".to_owned(), self.password.clone().unwrap()));
194 }
195 AUTH_METHOD_KEY_PAIR_FILE => {
196 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 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 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 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 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 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 "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(¶m.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 pub target_table_name: String,
666 pub database: String,
667 pub schema_name: String,
668 pub schema: Schema,
669
670 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 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 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(()); }
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 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 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 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}