Skip to main content

risingwave_connector/sink/
sqlserver.rs

1// Copyright 2024 RisingWave Labs
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::collections::{BTreeMap, HashMap};
16
17use anyhow::{Context, anyhow};
18use async_trait::async_trait;
19use phf::{Set, phf_set};
20use risingwave_common::array::{Op, RowRef, StreamChunk};
21use risingwave_common::catalog::Schema;
22use risingwave_common::row::{OwnedRow, Row};
23use risingwave_common::types::{DataType, Decimal};
24use serde::Deserialize;
25use serde_with::{DisplayFromStr, serde_as};
26use simd_json::prelude::ArrayTrait;
27use tiberius::numeric::Numeric;
28use tiberius::{AuthMethod, Client, ColumnData, Config, Query};
29use tokio::net::TcpStream;
30use tokio_util::compat::TokioAsyncWriteCompatExt;
31use with_options::WithOptions;
32
33use super::{
34    SINK_TYPE_APPEND_ONLY, SINK_TYPE_OPTION, SINK_TYPE_UPSERT, SinkError, SinkWriterMetrics,
35};
36use crate::enforce_secret::EnforceSecret;
37use crate::sink::writer::{LogSinkerOf, SinkWriter, SinkWriterExt};
38use crate::sink::{Result, Sink, SinkParam, SinkWriterParam};
39
40pub const SQLSERVER_SINK: &str = "sqlserver";
41
42fn default_max_batch_rows() -> usize {
43    1024
44}
45
46#[serde_as]
47#[derive(Clone, Debug, Deserialize, WithOptions)]
48pub struct SqlServerConfig {
49    #[serde(rename = "sqlserver.host")]
50    pub host: String,
51    #[serde(rename = "sqlserver.port")]
52    #[serde_as(as = "DisplayFromStr")]
53    pub port: u16,
54    #[serde(rename = "sqlserver.user")]
55    pub user: String,
56    #[serde(rename = "sqlserver.password")]
57    pub password: String,
58    #[serde(rename = "sqlserver.database")]
59    pub database: String,
60    #[serde(rename = "sqlserver.schema", default = "sql_server_default_schema")]
61    pub schema: String,
62    #[serde(rename = "sqlserver.table")]
63    pub table: String,
64    #[serde(
65        rename = "sqlserver.max_batch_rows",
66        default = "default_max_batch_rows"
67    )]
68    #[serde_as(as = "DisplayFromStr")]
69    pub max_batch_rows: usize,
70    pub r#type: String, // accept "append-only" or "upsert"
71
72    #[serde(flatten)]
73    pub unknown_fields: std::collections::HashMap<String, String>,
74}
75
76crate::impl_sink_unknown_fields!(SqlServerConfig);
77
78pub fn sql_server_default_schema() -> String {
79    "dbo".to_owned()
80}
81
82impl SqlServerConfig {
83    pub fn from_btreemap(properties: BTreeMap<String, String>) -> Result<Self> {
84        let config =
85            serde_json::from_value::<SqlServerConfig>(serde_json::to_value(properties).unwrap())
86                .map_err(|e| SinkError::Config(anyhow!(e)))?;
87        if config.r#type != SINK_TYPE_APPEND_ONLY && config.r#type != SINK_TYPE_UPSERT {
88            return Err(SinkError::Config(anyhow!(
89                "`{}` must be {}, or {}",
90                SINK_TYPE_OPTION,
91                SINK_TYPE_APPEND_ONLY,
92                SINK_TYPE_UPSERT
93            )));
94        }
95        Ok(config)
96    }
97
98    pub fn full_object_path(&self) -> String {
99        format!("[{}].[{}].[{}]", self.database, self.schema, self.table)
100    }
101}
102
103impl EnforceSecret for SqlServerConfig {
104    const ENFORCE_SECRET_PROPERTIES: Set<&'static str> = phf_set! {
105        "sqlserver.password"
106    };
107}
108#[derive(Debug)]
109pub struct SqlServerSink {
110    pub config: SqlServerConfig,
111    schema: Schema,
112    pk_indices: Vec<usize>,
113    is_append_only: bool,
114}
115
116struct SqlServerColumnMetadata {
117    name: String,
118    is_pk: bool,
119    data_type: String,
120}
121
122impl EnforceSecret for SqlServerSink {
123    fn enforce_secret<'a>(
124        prop_iter: impl Iterator<Item = &'a str>,
125    ) -> crate::sink::ConnectorResult<()> {
126        for prop in prop_iter {
127            SqlServerConfig::enforce_one(prop)?;
128        }
129        Ok(())
130    }
131}
132impl SqlServerSink {
133    pub fn new(
134        mut config: SqlServerConfig,
135        schema: Schema,
136        pk_indices: Vec<usize>,
137        is_append_only: bool,
138    ) -> Result<Self> {
139        // Rewrite config because tiberius allows a maximum of 2100 params in one query request.
140        const TIBERIUS_PARAM_MAX: usize = 2000;
141        let params_per_op = schema.fields().len();
142        let tiberius_max_batch_rows = if params_per_op == 0 {
143            config.max_batch_rows
144        } else {
145            ((TIBERIUS_PARAM_MAX as f64 / params_per_op as f64).floor()) as usize
146        };
147        if tiberius_max_batch_rows == 0 {
148            return Err(SinkError::SqlServer(anyhow!(format!(
149                "too many column {}",
150                params_per_op
151            ))));
152        }
153        config.max_batch_rows = std::cmp::min(config.max_batch_rows, tiberius_max_batch_rows);
154        Ok(Self {
155            config,
156            schema,
157            pk_indices,
158            is_append_only,
159        })
160    }
161}
162
163impl TryFrom<SinkParam> for SqlServerSink {
164    type Error = SinkError;
165
166    fn try_from(param: SinkParam) -> std::result::Result<Self, Self::Error> {
167        let schema = param.schema();
168        let pk_indices = param.downstream_pk_or_empty();
169        let config = SqlServerConfig::from_btreemap(param.properties)?;
170        SqlServerSink::new(config, schema, pk_indices, param.sink_type.is_append_only())
171    }
172}
173
174impl Sink for SqlServerSink {
175    type LogSinker = LogSinkerOf<SqlServerSinkWriter>;
176
177    const SINK_NAME: &'static str = SQLSERVER_SINK;
178
179    crate::impl_validate_sink_unknown_fields!();
180
181    async fn validate(&self) -> Result<()> {
182        risingwave_common::license::Feature::SqlServerSink
183            .check_available()
184            .map_err(|e| anyhow::anyhow!(e))?;
185
186        if !self.is_append_only && self.pk_indices.is_empty() {
187            return Err(SinkError::Config(anyhow!(
188                "Primary key not defined for upsert SQL Server sink (please define in `primary_key` field)"
189            )));
190        }
191
192        for f in self.schema.fields() {
193            check_data_type_compatibility(&f.data_type)?;
194        }
195
196        let mut sql_client = SqlServerClient::new(&self.config).await?;
197        validate_sql_server_write_permission(&mut sql_client, &self.config, self.is_append_only)
198            .await?;
199        let sql_server_table_metadata =
200            query_sql_server_table_metadata(&mut sql_client, &self.config).await?;
201        let sql_server_pk_count = sql_server_table_metadata
202            .iter()
203            .filter(|metadata| metadata.is_pk)
204            .count();
205        let sql_server_table_metadata = sql_server_table_metadata
206            .into_iter()
207            .map(|metadata| (metadata.name.clone(), metadata))
208            .collect::<HashMap<_, _>>();
209
210        // Validate Column name, Primary Key and data type.
211        for (idx, col) in self.schema.fields().iter().enumerate() {
212            let rw_is_pk = self.pk_indices.contains(&idx);
213            match sql_server_table_metadata.get(&normalize_sql_server_column_name(&col.name)) {
214                None => {
215                    return Err(SinkError::SqlServer(anyhow!(format!(
216                        "column {} not found in the downstream SQL Server table {}",
217                        col.name,
218                        self.config.full_object_path()
219                    ))));
220                }
221                Some(sql_server_col) => {
222                    validate_data_type_compatibility(
223                        &col.name,
224                        &col.data_type,
225                        &sql_server_col.data_type,
226                    )?;
227                    if self.is_append_only {
228                        continue;
229                    }
230                    if rw_is_pk && !sql_server_col.is_pk {
231                        return Err(SinkError::SqlServer(anyhow!(format!(
232                            "column {} specified in primary_key mismatches with the downstream SQL Server table {} PK",
233                            col.name,
234                            self.config.full_object_path(),
235                        ))));
236                    }
237                    if !rw_is_pk && sql_server_col.is_pk {
238                        return Err(SinkError::SqlServer(anyhow!(format!(
239                            "column {} unspecified in primary_key mismatches with the downstream SQL Server table {} PK",
240                            col.name,
241                            self.config.full_object_path(),
242                        ))));
243                    }
244                }
245            }
246        }
247
248        if !self.is_append_only && sql_server_pk_count != self.pk_indices.len() {
249            let sql_server_pk_columns = sql_server_table_metadata
250                .values()
251                .filter(|metadata| metadata.is_pk)
252                .map(|metadata| metadata.name.as_str())
253                .collect::<Vec<_>>()
254                .join(",");
255            let rw_pk_columns = self
256                .pk_indices
257                .iter()
258                .map(|idx| self.schema[*idx].name.as_str())
259                .collect::<Vec<_>>()
260                .join(",");
261            return Err(SinkError::SqlServer(anyhow!(format!(
262                "primary key does not match between RisingWave sink ({}: [{}]) and SQL Server table {} ({}: [{}])",
263                self.pk_indices.len(),
264                rw_pk_columns,
265                self.config.full_object_path(),
266                sql_server_pk_count,
267                sql_server_pk_columns,
268            ))));
269        }
270
271        Ok(())
272    }
273
274    async fn new_log_sinker(&self, writer_param: SinkWriterParam) -> Result<Self::LogSinker> {
275        Ok(SqlServerSinkWriter::new(
276            self.config.clone(),
277            self.schema.clone(),
278            self.pk_indices.clone(),
279            self.is_append_only,
280        )
281        .await?
282        .into_log_sinker(SinkWriterMetrics::new(&writer_param)))
283    }
284}
285
286enum SqlOp {
287    Insert(OwnedRow),
288    Merge(OwnedRow),
289    Delete(OwnedRow),
290}
291
292pub struct SqlServerSinkWriter {
293    config: SqlServerConfig,
294    schema: Schema,
295    pk_indices: Vec<usize>,
296    is_append_only: bool,
297    downstream_column_data_types: Vec<String>,
298    sql_client: SqlServerClient,
299    ops: Vec<SqlOp>,
300}
301
302impl SqlServerSinkWriter {
303    async fn new(
304        config: SqlServerConfig,
305        schema: Schema,
306        pk_indices: Vec<usize>,
307        is_append_only: bool,
308    ) -> Result<Self> {
309        let mut sql_client = SqlServerClient::new(&config).await?;
310        let downstream_column_data_types =
311            query_downstream_column_metadata(&mut sql_client, &config, &schema)
312                .await?
313                .into_iter()
314                .map(|metadata| metadata.data_type)
315                .collect();
316        let writer = Self {
317            config,
318            schema,
319            pk_indices,
320            is_append_only,
321            downstream_column_data_types,
322            sql_client,
323            ops: vec![],
324        };
325        Ok(writer)
326    }
327
328    async fn delete_one(&mut self, row: RowRef<'_>) -> Result<()> {
329        if self.ops.len() + 1 >= self.config.max_batch_rows {
330            self.flush().await?;
331        }
332        self.ops.push(SqlOp::Delete(row.into_owned_row()));
333        Ok(())
334    }
335
336    async fn upsert_one(&mut self, row: RowRef<'_>) -> Result<()> {
337        if self.ops.len() + 1 >= self.config.max_batch_rows {
338            self.flush().await?;
339        }
340        self.ops.push(SqlOp::Merge(row.into_owned_row()));
341        Ok(())
342    }
343
344    async fn insert_one(&mut self, row: RowRef<'_>) -> Result<()> {
345        if self.ops.len() + 1 >= self.config.max_batch_rows {
346            self.flush().await?;
347        }
348        self.ops.push(SqlOp::Insert(row.into_owned_row()));
349        Ok(())
350    }
351
352    async fn flush(&mut self) -> Result<()> {
353        use std::fmt::Write;
354        if self.ops.is_empty() {
355            return Ok(());
356        }
357        let mut query_str = String::new();
358        let col_num = self.schema.fields.len();
359        let mut next_param_id = 1;
360        let non_pk_col_indices = (0..col_num)
361            .filter(|idx| !self.pk_indices.contains(idx))
362            .collect::<Vec<usize>>();
363        let all_col_names = self
364            .schema
365            .fields
366            .iter()
367            .map(|f| format!("[{}]", f.name))
368            .collect::<Vec<_>>()
369            .join(",");
370        let all_source_col_names = self
371            .schema
372            .fields
373            .iter()
374            .map(|f| format!("[SOURCE].[{}]", f.name))
375            .collect::<Vec<_>>()
376            .join(",");
377        let pk_match = self
378            .pk_indices
379            .iter()
380            .map(|idx| {
381                format!(
382                    "[SOURCE].[{}]=[TARGET].[{}]",
383                    self.schema[*idx].name, self.schema[*idx].name
384                )
385            })
386            .collect::<Vec<_>>()
387            .join(" AND ");
388        let param_placeholders = |param_id: &mut usize| {
389            (0..col_num)
390                .map(|_| param_placeholder(param_id))
391                .collect::<Vec<_>>()
392                .join(",")
393        };
394        let set_all_source_col = non_pk_col_indices
395            .iter()
396            .map(|idx| {
397                format!(
398                    "[{}]=[SOURCE].[{}]",
399                    self.schema[*idx].name, self.schema[*idx].name
400                )
401            })
402            .collect::<Vec<_>>()
403            .join(",");
404        // TODO: avoid repeating the SQL
405        for op in &self.ops {
406            match op {
407                SqlOp::Insert(_) => {
408                    write!(
409                        &mut query_str,
410                        "INSERT INTO {} ({}) VALUES ({});",
411                        self.config.full_object_path(),
412                        all_col_names,
413                        param_placeholders(&mut next_param_id),
414                    )
415                    .unwrap();
416                }
417                SqlOp::Merge(_) => {
418                    write!(
419                        &mut query_str,
420                        r#"MERGE {} WITH (HOLDLOCK) AS [TARGET]
421                        USING (VALUES ({})) AS [SOURCE] ({})
422                        ON {}
423                        WHEN MATCHED THEN UPDATE SET {}
424                        WHEN NOT MATCHED THEN INSERT ({}) VALUES ({});"#,
425                        self.config.full_object_path(),
426                        param_placeholders(&mut next_param_id),
427                        all_col_names,
428                        pk_match,
429                        set_all_source_col,
430                        all_col_names,
431                        all_source_col_names,
432                    )
433                    .unwrap();
434                }
435                SqlOp::Delete(_) => {
436                    write!(
437                        &mut query_str,
438                        r#"DELETE FROM {} WHERE {};"#,
439                        self.config.full_object_path(),
440                        self.pk_indices
441                            .iter()
442                            .map(|idx| {
443                                let condition = format!(
444                                    "[{}]={}",
445                                    self.schema[*idx].name,
446                                    param_placeholder(&mut next_param_id)
447                                );
448                                condition
449                            })
450                            .collect::<Vec<_>>()
451                            .join(" AND "),
452                    )
453                    .unwrap();
454                }
455            }
456        }
457
458        let mut query = Query::new(query_str);
459        for op in self.ops.drain(..) {
460            match op {
461                SqlOp::Insert(row) => {
462                    bind_params(
463                        &mut query,
464                        row,
465                        &self.schema,
466                        &self.downstream_column_data_types,
467                        0..col_num,
468                    )?;
469                }
470                SqlOp::Merge(row) => {
471                    bind_params(
472                        &mut query,
473                        row,
474                        &self.schema,
475                        &self.downstream_column_data_types,
476                        0..col_num,
477                    )?;
478                }
479                SqlOp::Delete(row) => {
480                    bind_params(
481                        &mut query,
482                        row,
483                        &self.schema,
484                        &self.downstream_column_data_types,
485                        self.pk_indices.iter().copied(),
486                    )?;
487                }
488            }
489        }
490        query.execute(&mut self.sql_client.inner_client).await?;
491        Ok(())
492    }
493}
494
495#[async_trait]
496impl SinkWriter for SqlServerSinkWriter {
497    async fn begin_epoch(&mut self, _epoch: u64) -> Result<()> {
498        Ok(())
499    }
500
501    async fn write_batch(&mut self, chunk: StreamChunk) -> Result<()> {
502        for (op, row) in chunk.rows() {
503            match op {
504                Op::Insert => {
505                    if self.is_append_only {
506                        self.insert_one(row).await?;
507                    } else {
508                        self.upsert_one(row).await?;
509                    }
510                }
511                Op::UpdateInsert => {
512                    debug_assert!(!self.is_append_only);
513                    self.upsert_one(row).await?;
514                }
515                Op::Delete => {
516                    debug_assert!(!self.is_append_only);
517                    self.delete_one(row).await?;
518                }
519                Op::UpdateDelete => {}
520            }
521        }
522        Ok(())
523    }
524
525    async fn barrier(&mut self, is_checkpoint: bool) -> Result<Self::CommitMetadata> {
526        if is_checkpoint {
527            self.flush().await?;
528        }
529        Ok(())
530    }
531}
532
533#[derive(Debug)]
534pub struct SqlServerClient {
535    pub inner_client: Client<tokio_util::compat::Compat<TcpStream>>,
536}
537
538impl SqlServerClient {
539    async fn new(msconfig: &SqlServerConfig) -> Result<Self> {
540        let mut config = Config::new();
541        config.host(&msconfig.host);
542        config.port(msconfig.port);
543        config.authentication(AuthMethod::sql_server(&msconfig.user, &msconfig.password));
544        config.database(&msconfig.database);
545        config.trust_cert();
546        Self::new_with_config(config).await
547    }
548
549    pub async fn new_with_config(mut config: Config) -> Result<Self> {
550        let tcp = TcpStream::connect(config.get_addr())
551            .await
552            .context("failed to connect to sql server")
553            .map_err(SinkError::SqlServer)?;
554        tcp.set_nodelay(true)
555            .context("failed to setting nodelay when connecting to sql server")
556            .map_err(SinkError::SqlServer)?;
557
558        let client = match Client::connect(config.clone(), tcp.compat_write()).await {
559            // Connection successful.
560            Ok(client) => client,
561            // The server wants us to redirect to a different address
562            Err(tiberius::error::Error::Routing { host, port }) => {
563                config.host(&host);
564                config.port(port);
565                let tcp = TcpStream::connect(config.get_addr())
566                    .await
567                    .context("failed to connect to sql server after routing")
568                    .map_err(SinkError::SqlServer)?;
569                tcp.set_nodelay(true)
570                    .context(
571                        "failed to setting nodelay when connecting to sql server after routing",
572                    )
573                    .map_err(SinkError::SqlServer)?;
574                // we should not have more than one redirect, so we'll short-circuit here.
575                Client::connect(config, tcp.compat_write()).await?
576            }
577            Err(e) => return Err(e.into()),
578        };
579
580        Ok(Self {
581            inner_client: client,
582        })
583    }
584}
585
586async fn query_sql_server_table_metadata(
587    sql_client: &mut SqlServerClient,
588    config: &SqlServerConfig,
589) -> Result<Vec<SqlServerColumnMetadata>> {
590    let mut sql_server_table_metadata = Vec::new();
591    let query_table_metadata_error = || {
592        SinkError::SqlServer(anyhow!(format!(
593            "SQL Server table {} metadata error",
594            config.full_object_path()
595        )))
596    };
597    // Query primary-key membership through a subquery filtered by `pk.is_primary_key = 1`.
598    // A column can appear in both the primary-key index and secondary indexes, and a naive
599    // join from `sys.columns` to all `sys.index_columns` would emit extra index rows or mark
600    // secondary-index-only columns as PK columns. Keep the PK filter inside the subquery so
601    // each table column is returned once with `IsPk` set only by the primary-key index.
602    static QUERY_TABLE_METADATA: &str = r#"
603SELECT
604    col.name AS ColumnName,
605    CAST(CASE WHEN pk_col.column_id IS NULL THEN 0 ELSE 1 END AS int) AS IsPk,
606    typ.name AS DataType
607FROM
608    sys.columns col
609JOIN
610    sys.types typ ON typ.user_type_id = col.user_type_id
611LEFT JOIN
612    (
613        SELECT ic.object_id, ic.column_id
614        FROM sys.indexes pk
615        JOIN sys.index_columns ic ON ic.object_id = pk.object_id AND ic.index_id = pk.index_id
616        WHERE pk.is_primary_key = 1
617    ) pk_col ON pk_col.object_id = col.object_id AND pk_col.column_id = col.column_id
618WHERE
619    col.object_id = OBJECT_ID(@P1)
620ORDER BY
621    col.column_id;"#;
622    let rows = sql_client
623        .inner_client
624        .query(QUERY_TABLE_METADATA, &[&config.full_object_path()])
625        .await?
626        .into_results()
627        .await?;
628    for row in rows.into_iter().flatten() {
629        let mut iter = row.into_iter();
630        let ColumnData::String(Some(col_name)) =
631            iter.next().ok_or_else(query_table_metadata_error)?
632        else {
633            return Err(query_table_metadata_error());
634        };
635        let ColumnData::I32(Some(col_is_pk)) =
636            iter.next().ok_or_else(query_table_metadata_error)?
637        else {
638            return Err(query_table_metadata_error());
639        };
640        let ColumnData::String(Some(data_type)) =
641            iter.next().ok_or_else(query_table_metadata_error)?
642        else {
643            return Err(query_table_metadata_error());
644        };
645        sql_server_table_metadata.push(SqlServerColumnMetadata {
646            name: normalize_sql_server_column_name(&col_name),
647            is_pk: col_is_pk != 0,
648            data_type: data_type.into_owned(),
649        });
650    }
651    Ok(sql_server_table_metadata)
652}
653
654async fn validate_sql_server_write_permission(
655    sql_client: &mut SqlServerClient,
656    config: &SqlServerConfig,
657    is_append_only: bool,
658) -> Result<()> {
659    let permission_query_error = || {
660        SinkError::SqlServer(anyhow!(format!(
661            "SQL Server table {} permission metadata error",
662            config.full_object_path()
663        )))
664    };
665    static QUERY_WRITE_PERMISSION: &str = r#"
666SELECT
667    CAST(HAS_PERMS_BY_NAME(@P1, 'OBJECT', 'INSERT') AS int) AS CanInsert,
668    CAST(HAS_PERMS_BY_NAME(@P1, 'OBJECT', 'UPDATE') AS int) AS CanUpdate,
669    CAST(HAS_PERMS_BY_NAME(@P1, 'OBJECT', 'DELETE') AS int) AS CanDelete;"#;
670    let rows = sql_client
671        .inner_client
672        .query(QUERY_WRITE_PERMISSION, &[&config.full_object_path()])
673        .await?
674        .into_results()
675        .await?;
676    let mut rows = rows.into_iter().flatten();
677    let row = rows.next().ok_or_else(permission_query_error)?;
678    let mut iter = row.into_iter();
679    let ColumnData::I32(can_insert) = iter.next().ok_or_else(permission_query_error)? else {
680        return Err(permission_query_error());
681    };
682    let ColumnData::I32(can_update) = iter.next().ok_or_else(permission_query_error)? else {
683        return Err(permission_query_error());
684    };
685    let ColumnData::I32(can_delete) = iter.next().ok_or_else(permission_query_error)? else {
686        return Err(permission_query_error());
687    };
688
689    let missing_permissions = missing_sql_server_write_permissions(
690        is_append_only,
691        permission_is_granted(can_insert),
692        permission_is_granted(can_update),
693        permission_is_granted(can_delete),
694    );
695    if missing_permissions.is_empty() {
696        return Ok(());
697    }
698
699    Err(SinkError::SqlServer(anyhow!(format!(
700        "SQL Server user {} lacks required write permission(s) {} on table {}",
701        config.user,
702        missing_permissions.join(", "),
703        config.full_object_path()
704    ))))
705}
706
707fn permission_is_granted(permission_value: Option<i32>) -> bool {
708    permission_value == Some(1)
709}
710
711fn missing_sql_server_write_permissions(
712    is_append_only: bool,
713    can_insert: bool,
714    can_update: bool,
715    can_delete: bool,
716) -> Vec<&'static str> {
717    let mut missing_permissions = vec![];
718    if !can_insert {
719        missing_permissions.push("INSERT");
720    }
721    if !is_append_only {
722        if !can_update {
723            missing_permissions.push("UPDATE");
724        }
725        if !can_delete {
726            missing_permissions.push("DELETE");
727        }
728    }
729    missing_permissions
730}
731
732async fn query_downstream_column_metadata(
733    sql_client: &mut SqlServerClient,
734    config: &SqlServerConfig,
735    schema: &Schema,
736) -> Result<Vec<SqlServerColumnMetadata>> {
737    let sql_server_table_metadata = query_sql_server_table_metadata(sql_client, config)
738        .await?
739        .into_iter()
740        .map(|metadata| (metadata.name.clone(), metadata))
741        .collect::<HashMap<_, _>>();
742    schema
743        .fields()
744        .iter()
745        .map(|col| {
746            sql_server_table_metadata
747                .get(&normalize_sql_server_column_name(&col.name))
748                .map(|metadata| SqlServerColumnMetadata {
749                    name: metadata.name.clone(),
750                    is_pk: metadata.is_pk,
751                    data_type: metadata.data_type.clone(),
752                })
753                .ok_or_else(|| {
754                    SinkError::SqlServer(anyhow!(format!(
755                        "column {} not found in the downstream SQL Server table {}",
756                        col.name,
757                        config.full_object_path()
758                    )))
759                })
760        })
761        .collect()
762}
763
764fn param_placeholder(param_id: &mut usize) -> String {
765    let placeholder = format!("@P{}", *param_id);
766    *param_id += 1;
767    placeholder
768}
769
770fn bind_params(
771    query: &mut Query<'_>,
772    row: impl Row,
773    schema: &Schema,
774    downstream_column_data_types: &[String],
775    col_indices: impl Iterator<Item = usize>,
776) -> Result<()> {
777    use risingwave_common::types::ScalarRefImpl;
778    for col_idx in col_indices {
779        match row.datum_at(col_idx) {
780            Some(data_ref) => match data_ref {
781                ScalarRefImpl::Int16(v) => query.bind(v),
782                ScalarRefImpl::Int32(v) => query.bind(v),
783                ScalarRefImpl::Int64(v) => query.bind(v),
784                ScalarRefImpl::Float32(v) => query.bind(v.into_inner()),
785                ScalarRefImpl::Float64(v) => query.bind(v.into_inner()),
786                ScalarRefImpl::Utf8(v) => query.bind(v.to_owned()),
787                ScalarRefImpl::Bool(v) => query.bind(v),
788                ScalarRefImpl::Decimal(v) => match v {
789                    Decimal::Normalized(d) => {
790                        query.bind(decimal_to_sql(&d));
791                    }
792                    Decimal::NaN | Decimal::PositiveInf | Decimal::NegativeInf => {
793                        tracing::warn!(
794                            "Inf, -Inf, Nan in RisingWave decimal is converted into SQL Server null!"
795                        );
796                        query.bind(None as Option<Numeric>);
797                    }
798                },
799                ScalarRefImpl::Date(v) => query.bind(v.0),
800                ScalarRefImpl::Timestamp(v) => query.bind(v.0),
801                ScalarRefImpl::Timestamptz(v) => {
802                    let downstream_data_type = &downstream_column_data_types[col_idx];
803                    match downstream_data_type.as_str() {
804                        "bigint" | "int" | "smallint" | "tinyint" => {
805                            query.bind(v.timestamp_micros());
806                        }
807                        "datetimeoffset" => {
808                            query.bind(v.to_datetime_utc().fixed_offset());
809                        }
810                        "datetime" | "datetime2" | "smalldatetime" => {
811                            query.bind(v.to_datetime_utc().naive_utc());
812                        }
813                        _ => {
814                            return Err(unexpected_downstream_timestamptz_type(
815                                downstream_data_type,
816                            ));
817                        }
818                    };
819                }
820                ScalarRefImpl::Time(v) => query.bind(v.0),
821                ScalarRefImpl::Bytea(v) => query.bind(v.to_vec()),
822                ScalarRefImpl::Interval(_) => return Err(data_type_not_supported("Interval")),
823                ScalarRefImpl::Jsonb(_) => return Err(data_type_not_supported("Jsonb")),
824                ScalarRefImpl::Struct(_) => return Err(data_type_not_supported("Struct")),
825                ScalarRefImpl::List(_) => return Err(data_type_not_supported("List")),
826                ScalarRefImpl::Int256(_) => return Err(data_type_not_supported("Int256")),
827                ScalarRefImpl::Serial(_) => return Err(data_type_not_supported("Serial")),
828                ScalarRefImpl::Map(_) => return Err(data_type_not_supported("Map")),
829                ScalarRefImpl::Vector(_) => return Err(data_type_not_supported("Vector")),
830            },
831            None => match schema[col_idx].data_type {
832                DataType::Boolean => {
833                    query.bind(None as Option<bool>);
834                }
835                DataType::Int16 => {
836                    query.bind(None as Option<i16>);
837                }
838                DataType::Int32 => {
839                    query.bind(None as Option<i32>);
840                }
841                DataType::Int64 => {
842                    query.bind(None as Option<i64>);
843                }
844                DataType::Float32 => {
845                    query.bind(None as Option<f32>);
846                }
847                DataType::Float64 => {
848                    query.bind(None as Option<f64>);
849                }
850                DataType::Decimal => {
851                    query.bind(None as Option<Numeric>);
852                }
853                DataType::Date => {
854                    query.bind(None as Option<chrono::NaiveDate>);
855                }
856                DataType::Time => {
857                    query.bind(None as Option<chrono::NaiveTime>);
858                }
859                DataType::Timestamp => {
860                    query.bind(None as Option<chrono::NaiveDateTime>);
861                }
862                DataType::Timestamptz => {
863                    let downstream_data_type = &downstream_column_data_types[col_idx];
864                    match downstream_data_type.as_str() {
865                        "bigint" | "int" | "smallint" | "tinyint" => {
866                            query.bind(None as Option<i64>);
867                        }
868                        "datetimeoffset" => {
869                            query.bind(None as Option<chrono::DateTime<chrono::FixedOffset>>);
870                        }
871                        "datetime" | "datetime2" | "smalldatetime" => {
872                            query.bind(None as Option<chrono::NaiveDateTime>);
873                        }
874                        _ => {
875                            return Err(unexpected_downstream_timestamptz_type(
876                                downstream_data_type,
877                            ));
878                        }
879                    };
880                }
881                DataType::Varchar => {
882                    query.bind(None as Option<String>);
883                }
884                DataType::Bytea => {
885                    query.bind(None as Option<Vec<u8>>);
886                }
887                DataType::Interval => return Err(data_type_not_supported("Interval")),
888                DataType::Struct(_) => return Err(data_type_not_supported("Struct")),
889                DataType::List(_) => return Err(data_type_not_supported("List")),
890                DataType::Jsonb => return Err(data_type_not_supported("Jsonb")),
891                DataType::Serial => return Err(data_type_not_supported("Serial")),
892                DataType::Int256 => return Err(data_type_not_supported("Int256")),
893                DataType::Map(_) => return Err(data_type_not_supported("Map")),
894                DataType::Vector(_) => return Err(data_type_not_supported("Vector")),
895            },
896        };
897    }
898    Ok(())
899}
900
901fn data_type_not_supported(data_type_name: &str) -> SinkError {
902    SinkError::SqlServer(anyhow!(format!(
903        "{data_type_name} is not supported in SQL Server"
904    )))
905}
906
907fn unexpected_downstream_timestamptz_type(sql_server_data_type: &str) -> SinkError {
908    SinkError::SqlServer(anyhow!(format!(
909        "unexpected downstream SQL Server type {sql_server_data_type} for Timestamptz"
910    )))
911}
912
913fn check_data_type_compatibility(data_type: &DataType) -> Result<()> {
914    match data_type {
915        DataType::Boolean
916        | DataType::Int16
917        | DataType::Int32
918        | DataType::Int64
919        | DataType::Float32
920        | DataType::Float64
921        | DataType::Decimal
922        | DataType::Date
923        | DataType::Varchar
924        | DataType::Time
925        | DataType::Timestamp
926        | DataType::Timestamptz
927        | DataType::Bytea => Ok(()),
928        DataType::Interval => Err(data_type_not_supported("Interval")),
929        DataType::Struct(_) => Err(data_type_not_supported("Struct")),
930        DataType::List(_) => Err(data_type_not_supported("List")),
931        DataType::Jsonb => Err(data_type_not_supported("Jsonb")),
932        DataType::Serial => Err(data_type_not_supported("Serial")),
933        DataType::Int256 => Err(data_type_not_supported("Int256")),
934        DataType::Map(_) => Err(data_type_not_supported("Map")),
935        DataType::Vector(_) => Err(data_type_not_supported("Vector")),
936    }
937}
938
939fn normalize_sql_server_column_name(column_name: &str) -> String {
940    // SQL Server identifiers are usually case-insensitive depending on database collation.
941    // Match metadata by a case-insensitive key so validation follows that common behavior.
942    column_name.to_lowercase()
943}
944
945fn validate_data_type_compatibility(
946    column_name: &str,
947    rw_data_type: &DataType,
948    sql_server_data_type: &str,
949) -> Result<()> {
950    if sql_server_data_type_is_compatible(rw_data_type, sql_server_data_type) {
951        return Ok(());
952    }
953
954    Err(SinkError::SqlServer(anyhow!(format!(
955        "column {} data type {:?} is incompatible with downstream SQL Server type {}",
956        column_name, rw_data_type, sql_server_data_type
957    ))))
958}
959
960fn sql_server_data_type_is_compatible(rw_data_type: &DataType, sql_server_data_type: &str) -> bool {
961    match rw_data_type {
962        DataType::Boolean => sql_server_data_type == "bit",
963        DataType::Int16 => matches!(sql_server_data_type, "smallint" | "int" | "bigint"),
964        DataType::Int32 => matches!(sql_server_data_type, "int" | "bigint"),
965        DataType::Int64 => sql_server_data_type == "bigint",
966        DataType::Float32 => matches!(sql_server_data_type, "real" | "float"),
967        DataType::Float64 => sql_server_data_type == "float",
968        DataType::Decimal => matches!(sql_server_data_type, "decimal" | "numeric"),
969        DataType::Date => sql_server_data_type == "date",
970        DataType::Varchar => matches!(
971            sql_server_data_type,
972            "char" | "nchar" | "varchar" | "nvarchar" | "text" | "ntext"
973        ),
974        DataType::Time => sql_server_data_type == "time",
975        DataType::Timestamp => {
976            matches!(
977                sql_server_data_type,
978                "datetime" | "datetime2" | "smalldatetime"
979            )
980        }
981        DataType::Timestamptz => matches!(
982            sql_server_data_type,
983            "datetimeoffset"
984                | "datetime"
985                | "datetime2"
986                | "smalldatetime"
987                | "bigint"
988                | "int"
989                | "smallint"
990                | "tinyint"
991        ),
992        DataType::Bytea => matches!(sql_server_data_type, "binary" | "varbinary" | "image"),
993        DataType::Interval
994        | DataType::Struct(_)
995        | DataType::List(_)
996        | DataType::Jsonb
997        | DataType::Serial
998        | DataType::Int256
999        | DataType::Map(_)
1000        | DataType::Vector(_) => false,
1001    }
1002}
1003
1004/// The implementation is copied from tiberius crate.
1005fn decimal_to_sql(decimal: &rust_decimal::Decimal) -> Numeric {
1006    let unpacked = decimal.unpack();
1007
1008    let mut value = (((unpacked.hi as u128) << 64)
1009        + ((unpacked.mid as u128) << 32)
1010        + unpacked.lo as u128) as i128;
1011
1012    if decimal.is_sign_negative() {
1013        value = -value;
1014    }
1015
1016    Numeric::new_with_scale(value, decimal.scale() as u8)
1017}
1018
1019#[cfg(test)]
1020mod tests {
1021    use super::*;
1022
1023    #[test]
1024    fn test_normalize_sql_server_column_name() {
1025        assert_eq!(normalize_sql_server_column_name("EventDate"), "eventdate");
1026    }
1027
1028    #[test]
1029    fn test_sql_server_data_type_compatibility() {
1030        assert!(sql_server_data_type_is_compatible(
1031            &DataType::Int16,
1032            "smallint"
1033        ));
1034        assert!(sql_server_data_type_is_compatible(&DataType::Int16, "int"));
1035        assert!(!sql_server_data_type_is_compatible(
1036            &DataType::Int32,
1037            "smallint"
1038        ));
1039
1040        assert!(sql_server_data_type_is_compatible(
1041            &DataType::Timestamp,
1042            "datetime2"
1043        ));
1044        assert!(sql_server_data_type_is_compatible(
1045            &DataType::Timestamptz,
1046            "datetimeoffset"
1047        ));
1048        assert!(sql_server_data_type_is_compatible(
1049            &DataType::Timestamptz,
1050            "datetime2"
1051        ));
1052        assert!(sql_server_data_type_is_compatible(
1053            &DataType::Timestamptz,
1054            "bigint"
1055        ));
1056        assert!(sql_server_data_type_is_compatible(
1057            &DataType::Timestamptz,
1058            "int"
1059        ));
1060        assert!(!sql_server_data_type_is_compatible(
1061            &DataType::Timestamp,
1062            "datetimeoffset"
1063        ));
1064
1065        assert!(sql_server_data_type_is_compatible(
1066            &DataType::Varchar,
1067            "nvarchar"
1068        ));
1069        assert!(!sql_server_data_type_is_compatible(
1070            &DataType::Varchar,
1071            "uniqueidentifier"
1072        ));
1073    }
1074
1075    #[test]
1076    fn test_missing_sql_server_write_permissions() {
1077        assert_eq!(
1078            missing_sql_server_write_permissions(true, false, false, false),
1079            vec!["INSERT"]
1080        );
1081        assert!(missing_sql_server_write_permissions(true, true, false, false).is_empty());
1082        assert_eq!(
1083            missing_sql_server_write_permissions(false, false, false, false),
1084            vec!["INSERT", "UPDATE", "DELETE"]
1085        );
1086        assert_eq!(
1087            missing_sql_server_write_permissions(false, true, false, true),
1088            vec!["UPDATE"]
1089        );
1090        assert!(missing_sql_server_write_permissions(false, true, true, true).is_empty());
1091    }
1092}