1use 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, #[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 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 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 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 Ok(client) => client,
561 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 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 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 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
1004fn 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}