1use std::collections::{BTreeMap, HashMap, HashSet};
2
3use async_trait::async_trait;
4use chrono::{DateTime, FixedOffset, NaiveDate, NaiveDateTime, NaiveTime, Utc};
5use futures::TryStreamExt;
6use rustium_column_transform::ColumnTransformer;
7use rustium_config::{SnapshotConfig, SnapshotMode, SqlServerSourceConfig};
8use rustium_core::{
9 ChangeEvent, DataValue, Error, EventId, EventSchema, FieldSchema, Operation, RecordBoundary,
10 Result, RetryPolicy, Row, SignalAcknowledgement, SignalRecord, SourceConnector, SourceContext,
11 SourceMetadata, SourcePosition, SourceRecord, SqlServerPosition, TransactionMetadata,
12};
13use tiberius::{
14 AuthMethod, Client, ColumnData, Config as TdsConfig, EncryptionLevel, Row as TdsRow,
15};
16use tokio::net::TcpStream;
17use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt};
18use tracing::{info, warn};
19
20use crate::state::{
21 IncrementalSnapshotProgress, SqlServerKeyValue, decode_connector_state, encode_connector_state,
22};
23
24const CONNECTOR_VERSION: &str = env!("CARGO_PKG_VERSION");
25const LSN_SIZE: usize = 10;
26const RAW_DELETE: i32 = 1;
27const RAW_INSERT: i32 = 2;
28const RAW_UPDATE_BEFORE: i32 = 3;
29const RAW_UPDATE_AFTER: i32 = 4;
30const COMMIT_SERIAL: u64 = 5;
31const MAX_COMPLETED_SIGNAL_IDS: usize = 1_024;
32const WINDOW_OPEN_SIGNAL: &str = "snapshot-window-open";
33const WINDOW_CLOSE_SIGNAL: &str = "snapshot-window-close";
34
35type SqlClient = Client<Compat<TcpStream>>;
36
37#[derive(Debug, Clone)]
38struct CaptureTable {
39 schema: String,
40 table: String,
41 capture_instance: String,
42 source_object_id: i32,
43 event_schema: EventSchema,
44}
45
46impl CaptureTable {
47 fn key(&self) -> (String, String) {
48 (self.schema.clone(), self.table.clone())
49 }
50}
51
52#[derive(Debug, Clone)]
53struct CdcCursor {
54 commit_lsn: Vec<u8>,
55 change_lsn: Vec<u8>,
56 raw_operation: i32,
57}
58
59impl CdcCursor {
60 fn at_snapshot(commit_lsn: Vec<u8>) -> Self {
61 Self {
62 commit_lsn,
63 change_lsn: max_lsn_bytes(),
64 raw_operation: i32::try_from(COMMIT_SERIAL).unwrap_or(i32::MAX),
65 }
66 }
67
68 fn from_position(
69 position: &SqlServerPosition,
70 connector_state_boundary: bool,
71 ) -> Result<(Self, Option<SourcePosition>)> {
72 if position.snapshot || position.event_serial == COMMIT_SERIAL || connector_state_boundary {
73 return Ok((Self::at_snapshot(parse_lsn(&position.commit_lsn)?), None));
74 }
75 Ok((
76 Self {
77 commit_lsn: parse_lsn(&position.commit_lsn)?,
78 change_lsn: zero_lsn_bytes(),
79 raw_operation: 0,
80 },
81 Some(SourcePosition::SqlServer(position.clone())),
82 ))
83 }
84
85 fn commit_complete(&self) -> bool {
86 self.raw_operation == i32::try_from(COMMIT_SERIAL).unwrap_or(i32::MAX)
87 && self.change_lsn == max_lsn_bytes()
88 }
89}
90
91#[derive(Debug)]
92struct RawChange {
93 commit_lsn: Vec<u8>,
94 change_lsn: Vec<u8>,
95 raw_operation: i32,
96 capture_instance: String,
97 row: Row,
98 source_time: Option<DateTime<Utc>>,
99}
100
101#[derive(Debug)]
102struct ActiveTransaction {
103 commit_lsn: Vec<u8>,
104 source_time: Option<DateTime<Utc>>,
105 total_order: u64,
106 collection_order: HashMap<(String, String), u64>,
107}
108
109#[derive(Debug, Clone)]
110enum SnapshotSignal {
111 Execute {
112 id: String,
113 data_collections: Vec<String>,
114 additional_conditions: BTreeMap<String, String>,
115 },
116 Stop {
117 id: String,
118 data_collections: Vec<String>,
119 },
120 Pause {
121 id: String,
122 },
123 Resume {
124 id: String,
125 },
126 Unsupported {
127 id: String,
128 signal_type: String,
129 },
130}
131
132struct SqlServerIncrementalSnapshot {
133 progress: Option<IncrementalSnapshotProgress>,
134 opening_window_id: Option<String>,
135 window: Option<SqlServerIncrementalWindow>,
136 completed_signal_ids: Vec<String>,
137 event_serial: u64,
138 state_dirty: bool,
139}
140
141impl SqlServerIncrementalSnapshot {
142 fn new(
143 progress: Option<IncrementalSnapshotProgress>,
144 completed_signal_ids: Vec<String>,
145 ) -> Self {
146 Self {
147 progress,
148 opening_window_id: None,
149 window: None,
150 completed_signal_ids,
151 event_serial: 0,
152 state_dirty: false,
153 }
154 }
155
156 fn progress(&self) -> Option<&IncrementalSnapshotProgress> {
157 self.progress.as_ref()
158 }
159
160 fn completed_signal_ids(&self) -> &[String] {
161 &self.completed_signal_ids
162 }
163
164 fn is_active(&self) -> bool {
165 self.progress
166 .as_ref()
167 .is_some_and(|progress| !progress.paused)
168 && self.opening_window_id.is_none()
169 && self.window.is_none()
170 }
171
172 fn discard_window(&mut self) {
173 self.opening_window_id = None;
174 self.window = None;
175 }
176
177 fn state_dirty(&self) -> bool {
178 self.state_dirty
179 }
180
181 fn mark_checkpointed(&mut self) {
182 self.state_dirty = false;
183 }
184
185 fn remember_completed(&mut self, id: String) {
186 if let Some(index) = self
187 .completed_signal_ids
188 .iter()
189 .position(|candidate| candidate == &id)
190 {
191 self.completed_signal_ids.remove(index);
192 }
193 self.completed_signal_ids.push(id);
194 if self.completed_signal_ids.len() > MAX_COMPLETED_SIGNAL_IDS {
195 self.completed_signal_ids.remove(0);
196 }
197 }
198
199 fn parse_external_record(record: &SignalRecord) -> Result<SnapshotSignal> {
200 if record.id.trim().is_empty() || record.signal_type.trim().is_empty() {
201 return Err(Error::Source(
202 "SQL Server external signal requires non-empty id and type".into(),
203 ));
204 }
205 parse_snapshot_signal(&record.id, &record.signal_type, &record.data)
206 }
207
208 fn parse_row(row: &Row) -> Result<SnapshotSignal> {
209 let id = signal_text(row, "id")?;
210 let signal_type = signal_text(row, "type")?;
211 let data = signal_text(row, "data")?;
212 let value = serde_json::from_str::<serde_json::Value>(&data).map_err(|error| {
213 Error::Source(format!(
214 "SQL Server signal {id:?} has invalid JSON data: {error}"
215 ))
216 })?;
217 parse_snapshot_signal(&id, &signal_type, &value)
218 }
219
220 fn handle_signal(&mut self, signal: SnapshotSignal, source: &SqlServerSource) -> Result<()> {
221 match signal {
222 SnapshotSignal::Execute {
223 id,
224 data_collections,
225 additional_conditions,
226 } => {
227 if self.completed_signal_ids.contains(&id)
228 || self
229 .progress
230 .as_ref()
231 .is_some_and(|progress| progress.signal_id == id)
232 {
233 return Ok(());
234 }
235 if self.progress.is_some() {
236 tracing::warn!(%id, "SQL Server incremental snapshot is already active; execute signal ignored");
237 return Ok(());
238 }
239 source.signal_capture().ok_or_else(|| {
240 Error::Configuration(
241 "SQL Server incremental snapshots require a CDC-enabled signal.data.collection so Rustium can emit open/close watermarks"
242 .into(),
243 )
244 })?;
245 self.progress = Some(IncrementalSnapshotProgress {
246 signal_id: id,
247 data_collections: source.expand_data_collections(&data_collections)?,
248 additional_conditions,
249 current_collection: 0,
250 last_key: None,
251 maximum_key: None,
252 chunk_sequence: 1,
253 paused: false,
254 });
255 self.state_dirty = true;
256 Ok(())
257 }
258 SnapshotSignal::Stop {
259 id,
260 data_collections,
261 } => {
262 let Some(mut progress) = self.progress.take() else {
263 return Ok(());
264 };
265 if data_collections.is_empty() {
266 self.discard_window();
267 self.remember_completed(progress.signal_id);
268 self.state_dirty = true;
269 info!(%id, "SQL Server incremental snapshot stopped");
270 return Ok(());
271 }
272 let patterns = compile_collection_patterns(&data_collections)?;
273 let original = progress.data_collections.clone();
274 let current = original.get(progress.current_collection).cloned();
275 let retained_before = original
276 .iter()
277 .take(progress.current_collection)
278 .filter(|collection| !collection_matches_any(collection, &patterns))
279 .count();
280 progress
281 .data_collections
282 .retain(|collection| !collection_matches_any(collection, &patterns));
283 if progress.data_collections.len() == original.len() {
284 self.progress = Some(progress);
285 return Ok(());
286 }
287 self.discard_window();
288 if current
289 .as_ref()
290 .is_some_and(|collection| collection_matches_any(collection, &patterns))
291 {
292 progress.last_key = None;
293 progress.maximum_key = None;
294 }
295 progress.current_collection = retained_before;
296 if progress.current_collection >= progress.data_collections.len() {
297 self.remember_completed(progress.signal_id);
298 } else {
299 self.progress = Some(progress);
300 }
301 self.state_dirty = true;
302 info!(%id, "SQL Server incremental snapshot collections stopped");
303 Ok(())
304 }
305 SnapshotSignal::Pause { id } => {
306 if let Some(progress) = &mut self.progress
307 && !progress.paused
308 {
309 progress.paused = true;
310 self.state_dirty = true;
311 info!(%id, "SQL Server incremental snapshot paused");
312 }
313 Ok(())
314 }
315 SnapshotSignal::Resume { id } => {
316 if let Some(progress) = &mut self.progress
317 && progress.paused
318 {
319 progress.paused = false;
320 self.state_dirty = true;
321 info!(%id, "SQL Server incremental snapshot resumed");
322 }
323 Ok(())
324 }
325 SnapshotSignal::Unsupported { id, signal_type } => {
326 tracing::warn!(%id, %signal_type, "unsupported SQL Server runtime signal ignored");
327 Ok(())
328 }
329 }
330 }
331
332 async fn start_next_chunk(&mut self, source: &SqlServerSource) -> Result<()> {
333 if self.opening_window_id.is_some() || self.window.is_some() {
334 return Ok(());
335 }
336 let Some(progress) = self.progress.clone() else {
337 return Ok(());
338 };
339 if progress.paused {
340 return Ok(());
341 }
342 let window_id = uuid::Uuid::new_v4().to_string();
343 source
344 .emit_incremental_watermark(&format!("{window_id}-open"), WINDOW_OPEN_SIGNAL)
345 .await?;
346 self.opening_window_id = Some(window_id);
347 Ok(())
348 }
349
350 async fn open_window(&mut self, source: &SqlServerSource, window_id: String) -> Result<()> {
351 let Some(progress) = self.progress.clone() else {
352 return Ok(());
353 };
354 let collection = progress
355 .data_collections
356 .get(progress.current_collection)
357 .cloned()
358 .ok_or_else(|| {
359 Error::State(format!(
360 "SQL Server incremental snapshot collection index {} is outside {} collections",
361 progress.current_collection,
362 progress.data_collections.len()
363 ))
364 })?;
365 let capture = source.capture_for_collection(&collection).ok_or_else(|| {
366 Error::Source(format!(
367 "SQL Server incremental snapshot collection {collection:?} is not captured"
368 ))
369 })?;
370 let chunk = source.read_incremental_chunk(&capture, &progress).await?;
371 let row_count = chunk.rows.len();
372 let last_key = chunk.rows.last().map(|(_, key)| key.clone());
373 let remaining_keys = chunk
374 .rows
375 .iter()
376 .map(|(_, key)| key.clone())
377 .collect::<HashSet<_>>();
378 self.window = Some(SqlServerIncrementalWindow {
379 id: window_id.clone(),
380 collection,
381 capture,
382 rows: chunk.rows,
383 remaining_keys,
384 maximum_key: chunk.maximum_key,
385 last_key,
386 row_count,
387 close_commit_lsn: None,
388 });
389 source
390 .emit_incremental_watermark(&format!("{window_id}-close"), WINDOW_CLOSE_SIGNAL)
391 .await?;
392 Ok(())
393 }
394
395 async fn handle_internal_signal(
396 &mut self,
397 record: &SourceRecord,
398 row: &Row,
399 source: &SqlServerSource,
400 ) -> Result<bool> {
401 let signal_type = signal_text(row, "type")?;
402 if !matches!(
403 signal_type.as_str(),
404 WINDOW_OPEN_SIGNAL | WINDOW_CLOSE_SIGNAL
405 ) {
406 return Ok(false);
407 }
408 let id = signal_text(row, "id")?;
409 let opening_matches = self
410 .opening_window_id
411 .as_ref()
412 .is_some_and(|window_id| id == format!("{window_id}-open"));
413 let closing_matches = self
414 .window
415 .as_ref()
416 .is_some_and(|window| id == format!("{}-close", window.id));
417 if signal_type == WINDOW_OPEN_SIGNAL && opening_matches {
418 let window_id = self.opening_window_id.take().ok_or_else(|| {
419 Error::State("SQL Server incremental opening window disappeared".into())
420 })?;
421 self.open_window(source, window_id).await?;
422 } else if signal_type == WINDOW_CLOSE_SIGNAL && closing_matches {
423 let SourcePosition::SqlServer(position) = &record.position else {
424 return Err(Error::State(
425 "SQL Server watermark record has a non-SQL Server position".into(),
426 ));
427 };
428 let capture = self
429 .window
430 .as_ref()
431 .map(|window| window.capture.clone())
432 .ok_or_else(|| Error::State("SQL Server incremental window disappeared".into()))?;
433 source.validate_incremental_schema(&capture).await?;
434 if let Some(window) = &mut self.window {
435 window.close_commit_lsn = Some(position.commit_lsn.clone());
436 }
437 }
438 Ok(true)
439 }
440
441 fn observe_record(&mut self, record: &SourceRecord) -> Result<()> {
442 let Some(window) = &mut self.window else {
443 return Ok(());
444 };
445 if window.close_commit_lsn.is_some() {
446 return Ok(());
447 }
448 let Some(event) = &record.event else {
449 return Ok(());
450 };
451 let collection = event
452 .source
453 .schema
454 .as_deref()
455 .zip(event.source.table.as_deref())
456 .map(|(schema, table)| format!("{schema}.{table}"));
457 if collection.as_deref() != Some(window.collection.as_str()) {
458 return Ok(());
459 }
460 for row in [event.before.as_ref(), event.after.as_ref()]
461 .into_iter()
462 .flatten()
463 {
464 let key = sqlserver_key_from_row(row, &window.capture.event_schema)?;
465 window.remaining_keys.remove(&key);
466 }
467 Ok(())
468 }
469
470 fn closes_at(&self, record: &SourceRecord) -> bool {
471 let SourcePosition::SqlServer(position) = &record.position else {
472 return false;
473 };
474 record.boundary == RecordBoundary::TransactionCommit
475 && self.window.as_ref().is_some_and(|window| {
476 window.close_commit_lsn.as_deref() == Some(position.commit_lsn.as_str())
477 })
478 }
479
480 async fn finish_window(
481 &mut self,
482 source: &SqlServerSource,
483 base_position: &SourcePosition,
484 output: &tokio::sync::mpsc::Sender<Result<SourceRecord>>,
485 ) -> Result<SourcePosition> {
486 let window = self.window.take().ok_or_else(|| {
487 Error::State("SQL Server incremental snapshot has no pending window".into())
488 })?;
489 let progress = self.progress.clone().ok_or_else(|| {
490 Error::State("SQL Server incremental snapshot window has no progress".into())
491 })?;
492 for (row, key) in window.rows {
493 if !window.remaining_keys.contains(&key) {
494 continue;
495 }
496 self.event_serial = self.event_serial.saturating_add(1);
497 let position = sqlserver_incremental_position(base_position, self.event_serial)?;
498 let mut attributes = BTreeMap::new();
499 attributes.insert("rustium.snapshot.kind".into(), "incremental".into());
500 let event = ChangeEvent {
501 id: EventId::deterministic(
502 &source.connector_name,
503 source.database(),
504 &position,
505 &window.collection,
506 self.event_serial,
507 ),
508 source: SourceMetadata {
509 connector: "sqlserver".into(),
510 connector_name: source.connector_name.clone(),
511 database: source.database().into(),
512 schema: Some(window.capture.schema.clone()),
513 table: Some(window.capture.table.clone()),
514 snapshot: true,
515 version: CONNECTOR_VERSION.into(),
516 attributes,
517 },
518 position: position.clone(),
519 transaction: None,
520 operation: Operation::Read,
521 before: None,
522 after: Some(row),
523 schema: window.capture.event_schema.clone(),
524 source_time: None,
525 observed_time: Utc::now(),
526 };
527 output
528 .send(Ok(SourceRecord::data(event)))
529 .await
530 .map_err(|_| Error::Cancelled)?;
531 }
532
533 let mut next = progress;
534 if next.maximum_key.is_none() {
535 next.maximum_key.clone_from(&window.maximum_key);
536 }
537 next.last_key = window.last_key;
538 let collection_complete = next.maximum_key.is_none()
539 || next.last_key.as_ref() == next.maximum_key.as_ref()
540 || window.row_count < source.config.incremental_snapshot_chunk_size;
541 if collection_complete {
542 next.current_collection = next.current_collection.saturating_add(1);
543 next.last_key = None;
544 next.maximum_key = None;
545 }
546 next.chunk_sequence = next.chunk_sequence.saturating_add(1);
547 if next.current_collection >= next.data_collections.len() {
548 self.remember_completed(next.signal_id);
549 self.progress = None;
550 } else {
551 self.progress = Some(next);
552 }
553
554 self.event_serial = self.event_serial.saturating_add(1);
555 let position = sqlserver_incremental_position(base_position, self.event_serial)?;
556 output
557 .send(Ok(SourceRecord {
558 event: None,
559 position: position.clone(),
560 boundary: RecordBoundary::TransactionCommit,
561 connector_state: Some(encode_connector_state(
562 self.progress.as_ref(),
563 &self.completed_signal_ids,
564 )?),
565 signal_acknowledgements: Vec::new(),
566 }))
567 .await
568 .map_err(|_| Error::Cancelled)?;
569 self.mark_checkpointed();
570 Ok(position)
571 }
572}
573
574struct IncrementalChunk {
575 rows: Vec<(Row, Vec<SqlServerKeyValue>)>,
576 maximum_key: Option<Vec<SqlServerKeyValue>>,
577}
578
579struct SqlServerIncrementalWindow {
580 id: String,
581 collection: String,
582 capture: CaptureTable,
583 rows: Vec<(Row, Vec<SqlServerKeyValue>)>,
584 remaining_keys: HashSet<Vec<SqlServerKeyValue>>,
585 maximum_key: Option<Vec<SqlServerKeyValue>>,
586 last_key: Option<Vec<SqlServerKeyValue>>,
587 row_count: usize,
588 close_commit_lsn: Option<String>,
589}
590
591pub struct SqlServerSource {
592 connector_name: String,
593 config: SqlServerSourceConfig,
594 snapshot: SnapshotConfig,
595 captures: Vec<CaptureTable>,
596 retry_policy: RetryPolicy,
597 column_transformer: Option<ColumnTransformer>,
598}
599
600impl SqlServerSource {
601 #[must_use]
602 pub fn new(
603 connector_name: impl Into<String>,
604 config: SqlServerSourceConfig,
605 snapshot: SnapshotConfig,
606 ) -> Self {
607 Self {
608 connector_name: connector_name.into(),
609 config,
610 snapshot,
611 captures: Vec::new(),
612 retry_policy: RetryPolicy::default(),
613 column_transformer: None,
614 }
615 }
616
617 #[must_use]
618 pub fn with_retry_policy(mut self, retry_policy: RetryPolicy) -> Self {
619 self.retry_policy = retry_policy;
620 self
621 }
622
623 fn database(&self) -> &str {
624 &self.config.databases[0]
625 }
626
627 async fn validate_source(&mut self) -> Result<()> {
628 self.column_transformer =
629 Some(ColumnTransformer::new(&self.config.column_transformations)?);
630 let mut client = connect(&self.config, self.database()).await?;
631 let row = client
632 .simple_query(
633 "SELECT CAST(SERVERPROPERTY('ProductMajorVersion') AS int) AS major_version, \
634 CAST(is_cdc_enabled AS bit) AS is_cdc_enabled, \
635 CAST(snapshot_isolation_state AS int) AS snapshot_isolation_state \
636 FROM sys.databases WHERE name = DB_NAME()",
637 )
638 .await
639 .map_err(sqlserver_error)?
640 .into_row()
641 .await
642 .map_err(sqlserver_error)?
643 .ok_or_else(|| Error::Source("SQL Server did not return database metadata".into()))?;
644 let major_version = required::<i32>(&row, "major_version")?;
645 let cdc_enabled = required::<bool>(&row, "is_cdc_enabled")?;
646 let snapshot_isolation_state = required::<i32>(&row, "snapshot_isolation_state")?;
647 if major_version < 14 {
648 return Err(Error::Configuration(format!(
649 "SQL Server 2017 or newer is required; major version is {major_version}"
650 )));
651 }
652 if !cdc_enabled {
653 return Err(Error::Configuration(format!(
654 "CDC is not enabled for SQL Server database {:?}",
655 self.database()
656 )));
657 }
658 if self.config.snapshot_isolation_mode == "snapshot" && snapshot_isolation_state != 1 {
659 return Err(Error::Configuration(
660 "snapshot.isolation.mode=snapshot requires ALLOW_SNAPSHOT_ISOLATION ON".into(),
661 ));
662 }
663
664 let captures = discover_captures(&mut client, &self.config, &self.connector_name).await?;
665 let signal_table = signal_table_key(&self.config);
666 if captures
667 .iter()
668 .all(|capture| signal_table.as_ref() == Some(&capture.key()))
669 {
670 return Err(Error::Configuration(
671 "the SQL Server CDC capture instances and table filters select no tables".into(),
672 ));
673 }
674 if let Some(signal_table) = signal_table {
675 let capture = captures
676 .iter()
677 .find(|capture| capture.key() == signal_table)
678 .ok_or_else(|| {
679 Error::Configuration(format!(
680 "SQL Server signal table {}.{} is not CDC-enabled",
681 signal_table.0, signal_table.1
682 ))
683 })?;
684 validate_signal_schema(capture)?;
685 validate_signal_insert_permission(&mut client, capture).await?;
686 }
687 current_max_lsn(&mut client).await?;
688 client.close().await.map_err(sqlserver_error)?;
689 self.captures = captures;
690 Ok(())
691 }
692
693 async fn run_snapshot(
694 &self,
695 output: &tokio::sync::mpsc::Sender<Result<SourceRecord>>,
696 ) -> Result<Vec<u8>> {
697 let mut client = connect(&self.config, self.database()).await?;
698 client
699 .simple_query(snapshot_begin_sql(&self.config.snapshot_isolation_mode))
700 .await
701 .map_err(sqlserver_error)?
702 .into_results()
703 .await
704 .map_err(sqlserver_error)?;
705 let anchor = current_max_lsn(&mut client).await?;
706 let mut ordinal = 0_u64;
707 let mut captures = self.captures.clone();
708 let signal_table = signal_table_key(&self.config);
709 captures.retain(|capture| {
710 signal_table.as_ref() != Some(&capture.key())
711 && snapshot_includes(
712 &self.snapshot,
713 self.database(),
714 &capture.schema,
715 &capture.table,
716 )
717 });
718 captures.sort_by_key(CaptureTable::key);
719 for capture in &captures {
720 snapshot_table(
721 &mut client,
722 self.database(),
723 &self.connector_name,
724 capture,
725 &anchor,
726 &mut ordinal,
727 output,
728 self.column_transformer
729 .as_ref()
730 .expect("validated transformer"),
731 )
732 .await?;
733 }
734 client
735 .simple_query("COMMIT TRANSACTION")
736 .await
737 .map_err(sqlserver_error)?
738 .into_results()
739 .await
740 .map_err(sqlserver_error)?;
741 ordinal += 1;
742 output
743 .send(Ok(SourceRecord {
744 event: None,
745 position: sqlserver_position(
746 self.database(),
747 &anchor,
748 &zero_lsn_bytes(),
749 ordinal,
750 true,
751 ),
752 boundary: RecordBoundary::SnapshotComplete,
753 connector_state: None,
754 signal_acknowledgements: Vec::new(),
755 }))
756 .await
757 .map_err(|_| Error::Cancelled)?;
758 client.close().await.map_err(sqlserver_error)?;
759 Ok(anchor)
760 }
761
762 async fn current_anchor(&self) -> Result<Vec<u8>> {
763 let mut client = connect(&self.config, self.database()).await?;
764 let anchor = current_max_lsn(&mut client).await?;
765 client.close().await.map_err(sqlserver_error)?;
766 Ok(anchor)
767 }
768
769 fn signal_capture(&self) -> Option<CaptureTable> {
770 let signal = signal_table_key(&self.config)?;
771 self.captures
772 .iter()
773 .find(|capture| capture.key() == signal)
774 .cloned()
775 }
776
777 async fn emit_incremental_watermark(&self, id: &str, signal_type: &str) -> Result<()> {
778 let capture = self.signal_capture().ok_or_else(|| {
779 Error::Configuration(
780 "SQL Server incremental snapshots require a CDC-enabled signal.data.collection"
781 .into(),
782 )
783 })?;
784 let table = format!(
785 "{}.{}",
786 quote_identifier(&capture.schema),
787 quote_identifier(&capture.table)
788 );
789 let data = "{}";
790 let mut client = connect(&self.config, self.database()).await?;
791 client
792 .execute(
793 format!("INSERT INTO {table} ([id], [type], [data]) VALUES (@P1, @P2, @P3)"),
794 &[&id, &signal_type, &data],
795 )
796 .await
797 .map_err(sqlserver_error)?;
798 client.close().await.map_err(sqlserver_error)
799 }
800
801 async fn validate_incremental_schema(&self, capture: &CaptureTable) -> Result<()> {
802 let mut client = connect(&self.config, self.database()).await?;
803 let fields = discover_fields(&mut client, capture.source_object_id).await?;
804 client.close().await.map_err(sqlserver_error)?;
805 if fields != capture.event_schema.fields {
806 return Err(Error::Source(format!(
807 "SQL Server schema changed while incremental snapshot window for {}.{} was open",
808 capture.schema, capture.table
809 )));
810 }
811 Ok(())
812 }
813
814 fn capture_for_collection(&self, collection: &str) -> Option<CaptureTable> {
815 self.captures
816 .iter()
817 .find(|capture| format!("{}.{}", capture.schema, capture.table) == collection)
818 .cloned()
819 }
820
821 fn expand_data_collections(&self, patterns: &[String]) -> Result<Vec<String>> {
822 let patterns = compile_collection_patterns(patterns)?;
823 let signal_table = signal_table_key(&self.config);
824 let mut collections = self
825 .captures
826 .iter()
827 .filter(|capture| signal_table.as_ref() != Some(&capture.key()))
828 .filter_map(|capture| {
829 let short = format!("{}.{}", capture.schema, capture.table);
830 let qualified = format!("{}.{}", self.database(), short);
831 patterns
832 .iter()
833 .any(|pattern| pattern.is_match(&short) || pattern.is_match(&qualified))
834 .then_some(short)
835 })
836 .collect::<Vec<_>>();
837 collections.sort();
838 collections.dedup();
839 if collections.is_empty() {
840 return Err(Error::Source(
841 "SQL Server incremental snapshot patterns select no captured tables".into(),
842 ));
843 }
844 Ok(collections)
845 }
846
847 async fn read_incremental_chunk(
848 &self,
849 capture: &CaptureTable,
850 progress: &IncrementalSnapshotProgress,
851 ) -> Result<IncrementalChunk> {
852 let key_fields = capture
853 .event_schema
854 .fields
855 .iter()
856 .enumerate()
857 .filter(|(_, field)| field.primary_key)
858 .collect::<Vec<_>>();
859 if key_fields.is_empty() {
860 return Err(Error::Source(format!(
861 "SQL Server incremental snapshot table {}.{} requires a primary key",
862 capture.schema, capture.table
863 )));
864 }
865 let key_indices = key_fields
866 .iter()
867 .map(|(index, _)| *index)
868 .collect::<Vec<_>>();
869 let key_columns = key_fields
870 .iter()
871 .map(|(_, field)| format!("ct.{}", quote_identifier(&field.name)))
872 .collect::<Vec<_>>();
873 let collection = format!("{}.{}", capture.schema, capture.table);
874 let qualified_collection = format!("{}.{}", self.database(), collection);
875 let condition = progress
876 .additional_conditions
877 .iter()
878 .find_map(|(pattern, filter)| {
879 regex::Regex::new(&format!(r"^(?:{pattern})$"))
880 .ok()
881 .and_then(|pattern| {
882 (pattern.is_match(&collection) || pattern.is_match(&qualified_collection))
883 .then_some(filter)
884 })
885 });
886 let table = format!(
887 "{}.{}",
888 quote_identifier(&capture.schema),
889 quote_identifier(&capture.table)
890 );
891 let mut client = connect(&self.config, self.database()).await?;
892 let current_fields = discover_fields(&mut client, capture.source_object_id).await?;
893 if current_fields != capture.event_schema.fields {
894 return Err(Error::Source(format!(
895 "SQL Server schema changed before incremental snapshot query for {}.{}",
896 capture.schema, capture.table
897 )));
898 }
899 let maximum_key = match &progress.maximum_key {
900 Some(key) => Some(key.clone()),
901 None => {
902 let projections = key_fields
903 .iter()
904 .enumerate()
905 .map(|(projection_index, (_, field))| {
906 change_value_expression(projection_index, field)
907 })
908 .collect::<Vec<_>>()
909 .join(", ");
910 let where_clause = condition
911 .map(|condition| format!(" WHERE ({condition})"))
912 .unwrap_or_default();
913 let ordering = key_columns
914 .iter()
915 .map(|column| format!("{column} DESC"))
916 .collect::<Vec<_>>()
917 .join(", ");
918 let row = client
919 .simple_query(format!(
920 "SELECT TOP (1) {projections} FROM {table} AS ct{where_clause} ORDER BY {ordering}"
921 ))
922 .await
923 .map_err(sqlserver_error)?
924 .into_row()
925 .await
926 .map_err(sqlserver_error)?;
927 row.map(|row| sqlserver_key_from_projection(&row, &key_fields))
928 .transpose()?
929 }
930 };
931 let Some(maximum_key) = maximum_key else {
932 client.close().await.map_err(sqlserver_error)?;
933 return Ok(IncrementalChunk {
934 rows: Vec::new(),
935 maximum_key: None,
936 });
937 };
938
939 validate_key_width(&key_columns, &maximum_key, "maximum")?;
940 let mut predicates = Vec::new();
941 if let Some(condition) = condition {
942 predicates.push(format!("({condition})"));
943 }
944 if let Some(last_key) = &progress.last_key {
945 validate_key_width(&key_columns, last_key, "last")?;
946 predicates.push(sqlserver_key_predicate(&key_columns, last_key, true)?);
947 }
948 predicates.push(sqlserver_key_predicate(&key_columns, &maximum_key, false)?);
949 let projections = capture
950 .event_schema
951 .fields
952 .iter()
953 .enumerate()
954 .map(|(index, field)| change_value_expression(index, field))
955 .collect::<Vec<_>>()
956 .join(", ");
957 let query = format!(
958 "SELECT TOP ({}) {projections} FROM {table} AS ct WHERE {} ORDER BY {}",
959 self.config.incremental_snapshot_chunk_size,
960 predicates.join(" AND "),
961 key_columns.join(", ")
962 );
963 let rows = client
964 .simple_query(query)
965 .await
966 .map_err(sqlserver_error)?
967 .into_first_result()
968 .await
969 .map_err(sqlserver_error)?;
970 let rows = rows
971 .iter()
972 .map(|row| {
973 let mut converted = convert_tds_row(row, &capture.event_schema)?;
974 let key = key_indices
975 .iter()
976 .map(|index| {
977 let field = &capture.event_schema.fields[*index];
978 converted
979 .get(&field.name)
980 .ok_or_else(|| {
981 Error::Source(format!(
982 "SQL Server incremental row is missing key {:?}",
983 field.name
984 ))
985 })
986 .and_then(sqlserver_key_from_data_value)
987 })
988 .collect::<Result<Vec<_>>>()?;
989 transform_row(
990 self.column_transformer
991 .as_ref()
992 .expect("validated transformer"),
993 &mut converted,
994 self.database(),
995 capture,
996 );
997 Ok((converted, key))
998 })
999 .collect::<Result<Vec<_>>>()?;
1000 client.close().await.map_err(sqlserver_error)?;
1001 Ok(IncrementalChunk {
1002 rows,
1003 maximum_key: Some(maximum_key),
1004 })
1005 }
1006}
1007
1008#[async_trait]
1009impl SourceConnector for SqlServerSource {
1010 fn source_type(&self) -> &'static str {
1011 "sqlserver"
1012 }
1013
1014 async fn validate(&mut self) -> Result<()> {
1015 self.validate_source().await
1016 }
1017
1018 async fn run(&mut self, mut context: SourceContext) -> Result<()> {
1019 let column_transformer = self.column_transformer.clone().map_or_else(
1020 || ColumnTransformer::new(&self.config.column_transformations),
1021 Ok,
1022 )?;
1023 self.column_transformer = Some(column_transformer);
1024 let checkpoint = context.initial_checkpoint.clone();
1025 let snapshot_needed = match self.snapshot.mode {
1026 SnapshotMode::Never => false,
1027 SnapshotMode::Initial | SnapshotMode::WhenNeeded => checkpoint
1028 .as_ref()
1029 .is_none_or(|checkpoint| !checkpoint.snapshot_completed),
1030 };
1031 let mut checkpoint_position = checkpoint
1032 .as_ref()
1033 .map(|checkpoint| checkpoint.source_position.clone());
1034 if checkpoint_position
1035 .as_ref()
1036 .is_some_and(|position| !matches!(position, SourcePosition::SqlServer(_)))
1037 {
1038 return Err(Error::State(
1039 "SQL Server connector cannot resume from another source checkpoint".into(),
1040 ));
1041 }
1042
1043 let checkpoint_has_connector_state = !snapshot_needed
1044 && checkpoint
1045 .as_ref()
1046 .and_then(|checkpoint| checkpoint.connector_state.as_ref())
1047 .is_some();
1048 let (mut incremental_progress, mut completed_signal_ids) = if !snapshot_needed {
1049 checkpoint
1050 .as_ref()
1051 .and_then(|checkpoint| checkpoint.connector_state.as_ref())
1052 .map(decode_connector_state)
1053 .transpose()?
1054 .unwrap_or_default()
1055 } else {
1056 (None, Vec::new())
1057 };
1058
1059 let (mut cursor, mut resume_position) = if snapshot_needed {
1060 (
1061 CdcCursor::at_snapshot(self.run_snapshot(&context.output).await?),
1062 None,
1063 )
1064 } else if let Some(SourcePosition::SqlServer(position)) = &checkpoint_position {
1065 if position.database != self.database() {
1066 return Err(Error::State(format!(
1067 "SQL Server checkpoint belongs to database {:?}, not {:?}",
1068 position.database,
1069 self.database()
1070 )));
1071 }
1072 CdcCursor::from_position(position, checkpoint_has_connector_state)?
1073 } else {
1074 (CdcCursor::at_snapshot(self.current_anchor().await?), None)
1075 };
1076
1077 let mut client = connect(&self.config, self.database()).await?;
1078 if self.snapshot.mode == SnapshotMode::WhenNeeded && !snapshot_needed {
1079 if let Err(error) =
1080 validate_retention(&mut client, &self.captures, &cursor.commit_lsn).await
1081 {
1082 if matches!(&error, Error::State(_)) {
1083 warn!(
1084 connector = %self.connector_name,
1085 %error,
1086 "SQL Server checkpoint is older than CDC retention; taking a recovery snapshot"
1087 );
1088 client.close().await.map_err(sqlserver_error)?;
1089 let anchor = self.run_snapshot(&context.output).await?;
1090 cursor = CdcCursor::at_snapshot(anchor);
1091 resume_position = None;
1092 checkpoint_position = None;
1093 incremental_progress = None;
1094 completed_signal_ids.clear();
1095 client = connect(&self.config, self.database()).await?;
1096 } else {
1097 return Err(error);
1098 }
1099 }
1100 } else {
1101 validate_retention(&mut client, &self.captures, &cursor.commit_lsn).await?;
1102 }
1103 let mut state = StreamingState::new(cursor, resume_position);
1104 let mut incremental =
1105 SqlServerIncrementalSnapshot::new(incremental_progress, completed_signal_ids);
1106 let mut last_safe_position = checkpoint_position
1107 .as_ref()
1108 .filter(|position| {
1109 matches!(position, SourcePosition::SqlServer(position) if !position.snapshot)
1110 })
1111 .cloned()
1112 .unwrap_or_else(|| state.safe_position(self.database()));
1113 let mut interval = tokio::time::interval(self.config.poll_interval);
1114 interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
1115 let mut heartbeat = heartbeat_timer(self.config.heartbeat_interval);
1116 let mut file_signal_poll = file_signal_timer(&self.config);
1117 let mut incremental_tick = tokio::time::interval(std::time::Duration::from_millis(1));
1118 incremental_tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
1119 let mut heartbeat_connection = if self.config.heartbeat_interval.is_zero()
1120 || self.config.heartbeat_action_query.is_none()
1121 {
1122 None
1123 } else {
1124 Some(connect(&self.config, self.database()).await?)
1125 };
1126 info!(
1127 connector = %self.connector_name,
1128 database = %self.database(),
1129 commit_lsn = %format_lsn(&state.cursor.commit_lsn),
1130 max_retries = self.retry_policy.max_retries,
1131 initial_retry_delay_ms = self.retry_policy.initial_delay.as_millis(),
1132 max_retry_delay_ms = self.retry_policy.max_delay.as_millis(),
1133 "SQL Server CDC streaming started"
1134 );
1135
1136 loop {
1137 tokio::select! {
1138 _ = context.cancellation.cancelled() => {
1139 client.close().await.map_err(sqlserver_error)?;
1140 if let Some(connection) = heartbeat_connection {
1141 connection.close().await.map_err(sqlserver_error)?;
1142 }
1143 return Ok(());
1144 }
1145 changed = context.acknowledged.changed() => {
1146 if changed.is_err() {
1147 return Err(Error::Cancelled);
1148 }
1149 }
1150 _ = incremental_tick.tick(),
1151 if incremental.is_active()
1152 && state.transaction.is_none()
1153 && state.cursor.commit_complete() => {
1154 incremental.start_next_chunk(self).await?;
1155 }
1156 () = next_file_signal_poll(&mut file_signal_poll),
1157 if signal_channel_enabled(&self.config, "file")
1158 && state.transaction.is_none()
1159 && state.cursor.commit_complete() => {
1160 for line in crate::file_signal::read_and_clear(&self.config.signal_file).await? {
1161 let record = match serde_json::from_str::<SignalRecord>(&line) {
1162 Ok(record) => record,
1163 Err(error) => {
1164 tracing::warn!(%error, "invalid SQL Server file signal ignored");
1165 continue;
1166 }
1167 };
1168 let signal = match SqlServerIncrementalSnapshot::parse_external_record(&record) {
1169 Ok(signal) => signal,
1170 Err(error) => {
1171 tracing::warn!(%error, "invalid SQL Server file signal ignored");
1172 continue;
1173 }
1174 };
1175 incremental.handle_signal(signal, self)?;
1176 }
1177 if incremental.state_dirty() {
1178 emit_incremental_checkpoint(
1179 &context.output,
1180 incremental.progress(),
1181 incremental.completed_signal_ids(),
1182 &last_safe_position,
1183 None,
1184 ).await?;
1185 incremental.mark_checkpointed();
1186 }
1187 }
1188 delivery = context.signals.recv(),
1189 if (signal_channel_enabled(&self.config, "in-process")
1190 || signal_channel_enabled(&self.config, "kafka"))
1191 && state.transaction.is_none()
1192 && state.cursor.commit_complete() => {
1193 let delivery = delivery.ok_or_else(|| {
1194 Error::Source("SQL Server runtime signal channel closed".into())
1195 })?;
1196 let signal = match SqlServerIncrementalSnapshot::parse_external_record(delivery.record()) {
1197 Ok(signal) => signal,
1198 Err(error) => {
1199 tracing::warn!(%error, "invalid SQL Server runtime signal ignored");
1200 delivery.acknowledge();
1201 continue;
1202 }
1203 };
1204 incremental.handle_signal(signal, self)?;
1205 emit_incremental_checkpoint(
1206 &context.output,
1207 incremental.progress(),
1208 incremental.completed_signal_ids(),
1209 &last_safe_position,
1210 delivery.into_acknowledgement(),
1211 ).await?;
1212 incremental.mark_checkpointed();
1213 }
1214 _ = interval.tick() => {
1215 let Some(records) = poll_change_batch_with_retry(
1216 self,
1217 &mut client,
1218 &mut state,
1219 &context.cancellation,
1220 ).await? else {
1221 continue;
1222 };
1223 for mut record in records {
1224 if let Some(signal_row) = source_signal_row(&record, &self.config) {
1225 if let Some(signal_row) = signal_row {
1226 if incremental
1227 .handle_internal_signal(&record, signal_row, self)
1228 .await?
1229 {
1230 continue;
1231 }
1232 if signal_channel_enabled(&self.config, "source") {
1233 let signal =
1234 SqlServerIncrementalSnapshot::parse_row(signal_row)?;
1235 incremental.handle_signal(signal, self)?;
1236 }
1237 }
1238 continue;
1239 }
1240 incremental.observe_record(&record)?;
1241 if incremental.closes_at(&record) {
1242 let base = record.position.clone();
1243 last_safe_position = incremental
1244 .finish_window(self, &base, &context.output)
1245 .await?;
1246 continue;
1247 }
1248 if record.boundary == RecordBoundary::TransactionCommit
1249 && incremental.state_dirty()
1250 {
1251 record.connector_state = Some(encode_connector_state(
1252 incremental.progress(),
1253 incremental.completed_signal_ids(),
1254 )?);
1255 }
1256 let checkpointed = record.connector_state.is_some();
1257 let position = record.position.clone();
1258 context.output.send(Ok(record)).await.map_err(|_| Error::Cancelled)?;
1259 if position.is_after(&last_safe_position)
1260 || position == last_safe_position
1261 {
1262 last_safe_position = position;
1263 }
1264 if checkpointed {
1265 incremental.mark_checkpointed();
1266 }
1267 }
1268 }
1269 () = next_heartbeat(&mut heartbeat) => {
1270 if !state.cursor.commit_complete() {
1271 continue;
1272 }
1273 if let Some(query) = self.config.heartbeat_action_query.clone() {
1274 let connection = heartbeat_connection.take().ok_or_else(|| {
1275 Error::Source("SQL Server heartbeat action connection is unavailable".into())
1276 })?;
1277 heartbeat_connection = Some(execute_heartbeat_action(connection, query).await?);
1278 }
1279 context.output.send(Ok(sqlserver_heartbeat_record(
1280 &self.connector_name,
1281 self.database(),
1282 last_safe_position.clone(),
1283 ))).await.map_err(|_| Error::Cancelled)?;
1284 }
1285 }
1286 }
1287 }
1288}
1289
1290async fn poll_change_batch_with_retry(
1291 source: &SqlServerSource,
1292 client: &mut SqlClient,
1293 state: &mut StreamingState,
1294 cancellation: &tokio_util::sync::CancellationToken,
1295) -> Result<Option<Vec<SourceRecord>>> {
1296 let mut retries = 0_u64;
1297 let mut delay = source.retry_policy.initial_delay;
1298 let mut reconnect = false;
1299 loop {
1300 if reconnect {
1301 match connect(&source.config, source.database()).await {
1302 Ok(reconnected) => {
1303 *client = reconnected;
1304 info!(
1305 connector = %source.connector_name,
1306 database = %source.database(),
1307 retries,
1308 "SQL Server CDC polling connection recovered"
1309 );
1310 }
1311 Err(error @ Error::RetryableSource(_)) => {
1312 if !wait_for_sqlserver_retry(source, retries, &mut delay, &error, cancellation)
1313 .await?
1314 {
1315 return Err(error);
1316 }
1317 retries += 1;
1318 continue;
1319 }
1320 Err(error) => return Err(error),
1321 }
1322 }
1323
1324 match poll_change_batch_once(source, client, state).await {
1325 Ok(records) => return Ok(records),
1326 Err(error @ Error::RetryableSource(_)) => {
1327 if !wait_for_sqlserver_retry(source, retries, &mut delay, &error, cancellation)
1328 .await?
1329 {
1330 return Err(error);
1331 }
1332 retries += 1;
1333 reconnect = true;
1334 }
1335 Err(error) => return Err(error),
1336 }
1337 }
1338}
1339
1340async fn wait_for_sqlserver_retry(
1341 source: &SqlServerSource,
1342 retries: u64,
1343 delay: &mut std::time::Duration,
1344 error: &Error,
1345 cancellation: &tokio_util::sync::CancellationToken,
1346) -> Result<bool> {
1347 if !sqlserver_retry_allowed(&source.retry_policy, retries) {
1348 return Ok(false);
1349 }
1350 warn!(
1351 connector = %source.connector_name,
1352 database = %source.database(),
1353 retry = retries + 1,
1354 max_retries = source.retry_policy.max_retries,
1355 delay_ms = delay.as_millis(),
1356 %error,
1357 "retryable SQL Server CDC polling failure; reconnecting"
1358 );
1359 tokio::select! {
1360 () = cancellation.cancelled() => return Err(Error::Cancelled),
1361 () = tokio::time::sleep(*delay) => {}
1362 }
1363 *delay = delay.saturating_mul(2).min(source.retry_policy.max_delay);
1364 Ok(true)
1365}
1366
1367fn sqlserver_retry_allowed(policy: &RetryPolicy, retries: u64) -> bool {
1368 policy.max_retries < 0 || retries < policy.max_retries as u64
1369}
1370
1371async fn poll_change_batch_once(
1372 source: &SqlServerSource,
1373 client: &mut SqlClient,
1374 state: &mut StreamingState,
1375) -> Result<Option<Vec<SourceRecord>>> {
1376 let max_lsn = current_max_lsn(client).await?;
1377 if max_lsn < state.cursor.commit_lsn
1378 || (max_lsn == state.cursor.commit_lsn && state.cursor.commit_complete())
1379 {
1380 return Ok(None);
1381 }
1382 validate_retention(client, &source.captures, &state.cursor.commit_lsn).await?;
1383 read_change_batch(
1384 client,
1385 source.database(),
1386 &source.connector_name,
1387 &source.captures,
1388 state,
1389 &max_lsn,
1390 source.config.streaming_fetch_size,
1391 source
1392 .column_transformer
1393 .as_ref()
1394 .expect("validated transformer"),
1395 signal_table_key(&source.config).as_ref(),
1396 )
1397 .await
1398 .map(Some)
1399}
1400
1401fn heartbeat_timer(interval: std::time::Duration) -> Option<tokio::time::Interval> {
1402 if interval.is_zero() {
1403 return None;
1404 }
1405 let mut timer = tokio::time::interval_at(tokio::time::Instant::now() + interval, interval);
1406 timer.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
1407 Some(timer)
1408}
1409
1410async fn next_heartbeat(timer: &mut Option<tokio::time::Interval>) {
1411 match timer {
1412 Some(timer) => {
1413 timer.tick().await;
1414 }
1415 None => std::future::pending::<()>().await,
1416 }
1417}
1418
1419fn signal_channel_enabled(config: &SqlServerSourceConfig, channel: &str) -> bool {
1420 config
1421 .signal_enabled_channels
1422 .iter()
1423 .any(|enabled| enabled == channel)
1424}
1425
1426fn signal_table_key(config: &SqlServerSourceConfig) -> Option<(String, String)> {
1427 let parts = config
1428 .signal_data_collection
1429 .as_deref()?
1430 .split('.')
1431 .collect::<Vec<_>>();
1432 match parts.as_slice() {
1433 [schema, table] | [_, schema, table] => Some(((*schema).into(), (*table).into())),
1434 _ => None,
1435 }
1436}
1437
1438fn transform_event(
1439 transformer: &ColumnTransformer,
1440 event: &mut ChangeEvent,
1441 database: &str,
1442 capture: &CaptureTable,
1443) {
1444 transformer.transform_event_with_qualified_tables(
1445 event,
1446 &qualified_table_candidates(database, capture),
1447 &HashMap::new(),
1448 );
1449}
1450
1451fn transform_row(
1452 transformer: &ColumnTransformer,
1453 row: &mut Row,
1454 database: &str,
1455 capture: &CaptureTable,
1456) {
1457 transformer.transform_row_with_qualified_tables(
1458 row,
1459 &qualified_table_candidates(database, capture),
1460 &capture.event_schema,
1461 &HashMap::new(),
1462 );
1463}
1464
1465fn qualified_table_candidates(database: &str, capture: &CaptureTable) -> [String; 2] {
1466 [
1467 format!("{database}.{}.{}", capture.schema, capture.table),
1468 format!("{}.{}", capture.schema, capture.table),
1469 ]
1470}
1471
1472fn snapshot_includes(snapshot: &SnapshotConfig, database: &str, schema: &str, table: &str) -> bool {
1473 snapshot.includes_collection(&format!("{database}.{schema}.{table}"))
1474}
1475
1476fn validate_signal_schema(capture: &CaptureTable) -> Result<()> {
1477 let expected = ["id", "type", "data"];
1478 if capture.event_schema.fields.len() != expected.len()
1479 || capture
1480 .event_schema
1481 .fields
1482 .iter()
1483 .zip(expected)
1484 .any(|(field, expected)| {
1485 field.name != expected
1486 || !matches!(
1487 base_type(&field.type_name),
1488 "char" | "varchar" | "nchar" | "nvarchar" | "text" | "ntext"
1489 )
1490 })
1491 {
1492 return Err(Error::Configuration(format!(
1493 "SQL Server signal table {}.{} must contain exactly text-compatible id, type, and data columns in that order",
1494 capture.schema, capture.table
1495 )));
1496 }
1497 let minimum_lengths = [
1498 uuid::Uuid::nil().to_string().len() + "-close".len(),
1499 WINDOW_OPEN_SIGNAL.len().max(WINDOW_CLOSE_SIGNAL.len()),
1500 2,
1501 ];
1502 for (field, minimum) in capture.event_schema.fields.iter().zip(minimum_lengths) {
1503 if signal_text_capacity(&field.type_name).is_some_and(|capacity| capacity < minimum) {
1504 return Err(Error::Configuration(format!(
1505 "SQL Server signal table column {:?} must hold at least {minimum} characters for incremental snapshot watermarks; found {}",
1506 field.name, field.type_name
1507 )));
1508 }
1509 }
1510 Ok(())
1511}
1512
1513fn signal_text_capacity(type_name: &str) -> Option<usize> {
1514 let normalized = type_name.trim().to_ascii_lowercase();
1515 if matches!(base_type(&normalized), "text" | "ntext") || normalized.ends_with("(max)") {
1516 return None;
1517 }
1518 normalized
1519 .split_once('(')
1520 .and_then(|(_, length)| length.strip_suffix(')'))
1521 .and_then(|length| length.parse().ok())
1522}
1523
1524async fn validate_signal_insert_permission(
1525 client: &mut SqlClient,
1526 capture: &CaptureTable,
1527) -> Result<()> {
1528 let object_name = format!("{}.{}", capture.schema, capture.table);
1529 let row = client
1530 .query(
1531 "SELECT CAST(HAS_PERMS_BY_NAME(@P1, 'OBJECT', 'INSERT') AS int) AS can_insert",
1532 &[&object_name],
1533 )
1534 .await
1535 .map_err(sqlserver_error)?
1536 .into_row()
1537 .await
1538 .map_err(sqlserver_error)?
1539 .ok_or_else(|| {
1540 Error::Source("SQL Server did not return signal-table permissions".into())
1541 })?;
1542 if required::<i32>(&row, "can_insert")? != 1 {
1543 return Err(Error::Configuration(format!(
1544 "SQL Server connector user requires INSERT on signal table {}.{} for incremental snapshot watermarks",
1545 capture.schema, capture.table
1546 )));
1547 }
1548 Ok(())
1549}
1550
1551fn source_signal_row<'a>(
1552 record: &'a SourceRecord,
1553 config: &SqlServerSourceConfig,
1554) -> Option<Option<&'a Row>> {
1555 let signal = signal_table_key(config)?;
1556 let event = record.event.as_ref()?;
1557 if event.source.schema.as_deref() != Some(signal.0.as_str())
1558 || event.source.table.as_deref() != Some(signal.1.as_str())
1559 {
1560 return None;
1561 }
1562 Some(
1563 (event.operation == Operation::Create)
1564 .then_some(event.after.as_ref())
1565 .flatten(),
1566 )
1567}
1568
1569fn file_signal_timer(config: &SqlServerSourceConfig) -> Option<tokio::time::Interval> {
1570 if !signal_channel_enabled(config, "file") {
1571 return None;
1572 }
1573 let mut timer = tokio::time::interval_at(
1574 tokio::time::Instant::now() + config.signal_poll_interval,
1575 config.signal_poll_interval,
1576 );
1577 timer.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
1578 Some(timer)
1579}
1580
1581async fn next_file_signal_poll(timer: &mut Option<tokio::time::Interval>) {
1582 match timer {
1583 Some(timer) => {
1584 timer.tick().await;
1585 }
1586 None => std::future::pending::<()>().await,
1587 }
1588}
1589
1590fn signal_text(row: &Row, name: &str) -> Result<String> {
1591 match row.get(name) {
1592 Some(DataValue::String(value)) => Ok(value.clone()),
1593 Some(DataValue::Json(value)) => Ok(value.to_string()),
1594 Some(DataValue::Bytes(value)) => String::from_utf8(value.clone()).map_err(|error| {
1595 Error::Source(format!(
1596 "SQL Server signal column {name} is not UTF-8: {error}"
1597 ))
1598 }),
1599 Some(value) => Ok(value.to_json("__rustium_unavailable").to_string()),
1600 None => Err(Error::Source(format!(
1601 "SQL Server signal table is missing column {name:?}"
1602 ))),
1603 }
1604}
1605
1606fn parse_snapshot_signal(
1607 id: &str,
1608 signal_type: &str,
1609 data: &serde_json::Value,
1610) -> Result<SnapshotSignal> {
1611 match signal_type {
1612 "execute-snapshot" => {
1613 let snapshot_type = data.get("type").and_then(serde_json::Value::as_str);
1614 if snapshot_type.is_some_and(|kind| !kind.eq_ignore_ascii_case("incremental")) {
1615 return Err(Error::Source(format!(
1616 "SQL Server execute-snapshot signal {id:?} supports only type=incremental"
1617 )));
1618 }
1619 let collections = data_collections(data);
1620 if collections.is_empty() {
1621 return Err(Error::Source(format!(
1622 "SQL Server execute-snapshot signal {id:?} has no data-collections"
1623 )));
1624 }
1625 compile_collection_patterns(&collections)?;
1626 let mut additional_conditions = BTreeMap::new();
1627 if let Some(values) = data
1628 .get("additional-conditions")
1629 .and_then(serde_json::Value::as_array)
1630 {
1631 for value in values {
1632 let collection = value
1633 .get("data-collection")
1634 .and_then(serde_json::Value::as_str)
1635 .ok_or_else(|| {
1636 Error::Source(format!(
1637 "SQL Server execute-snapshot signal {id:?} has an invalid additional-condition"
1638 ))
1639 })?;
1640 let filter = value
1641 .get("filter")
1642 .and_then(serde_json::Value::as_str)
1643 .ok_or_else(|| {
1644 Error::Source(format!(
1645 "SQL Server execute-snapshot signal {id:?} has an invalid additional-condition"
1646 ))
1647 })?;
1648 compile_collection_patterns(&[collection.into()])?;
1649 additional_conditions.insert(collection.into(), filter.into());
1650 }
1651 }
1652 Ok(SnapshotSignal::Execute {
1653 id: id.into(),
1654 data_collections: collections,
1655 additional_conditions,
1656 })
1657 }
1658 "stop-snapshot" => Ok(SnapshotSignal::Stop {
1659 id: id.into(),
1660 data_collections: data_collections(data),
1661 }),
1662 "pause-snapshot" => Ok(SnapshotSignal::Pause { id: id.into() }),
1663 "resume-snapshot" => Ok(SnapshotSignal::Resume { id: id.into() }),
1664 other => Ok(SnapshotSignal::Unsupported {
1665 id: id.into(),
1666 signal_type: other.into(),
1667 }),
1668 }
1669}
1670
1671fn data_collections(data: &serde_json::Value) -> Vec<String> {
1672 data.get("data-collections")
1673 .and_then(serde_json::Value::as_array)
1674 .map(|values| {
1675 values
1676 .iter()
1677 .filter_map(serde_json::Value::as_str)
1678 .map(str::to_owned)
1679 .collect()
1680 })
1681 .unwrap_or_default()
1682}
1683
1684fn compile_collection_patterns(patterns: &[String]) -> Result<Vec<regex::Regex>> {
1685 patterns
1686 .iter()
1687 .map(|pattern| {
1688 regex::Regex::new(&format!(r"^(?:{pattern})$")).map_err(|error| {
1689 Error::Source(format!(
1690 "invalid SQL Server snapshot collection pattern {pattern:?}: {error}"
1691 ))
1692 })
1693 })
1694 .collect()
1695}
1696
1697fn collection_matches_any(collection: &str, patterns: &[regex::Regex]) -> bool {
1698 patterns.iter().any(|pattern| pattern.is_match(collection))
1699}
1700
1701fn validate_key_width(columns: &[String], key: &[SqlServerKeyValue], boundary: &str) -> Result<()> {
1702 if columns.len() != key.len() {
1703 return Err(Error::State(format!(
1704 "SQL Server incremental snapshot {boundary} key has {} values for {} primary-key columns",
1705 key.len(),
1706 columns.len()
1707 )));
1708 }
1709 Ok(())
1710}
1711
1712fn sqlserver_key_from_projection(
1713 row: &TdsRow,
1714 key_fields: &[(usize, &FieldSchema)],
1715) -> Result<Vec<SqlServerKeyValue>> {
1716 key_fields
1717 .iter()
1718 .enumerate()
1719 .map(|(projection_index, (_, field))| {
1720 convert_tds_value(row, projection_index, &field.type_name)
1721 .and_then(|value| sqlserver_key_from_data_value(&value))
1722 })
1723 .collect()
1724}
1725
1726fn sqlserver_key_from_row(row: &Row, schema: &EventSchema) -> Result<Vec<SqlServerKeyValue>> {
1727 let key = schema
1728 .fields
1729 .iter()
1730 .filter(|field| field.primary_key)
1731 .map(|field| {
1732 row.get(&field.name)
1733 .ok_or_else(|| {
1734 Error::Source(format!(
1735 "SQL Server CDC row is missing primary-key column {:?}",
1736 field.name
1737 ))
1738 })
1739 .and_then(sqlserver_key_from_data_value)
1740 })
1741 .collect::<Result<Vec<_>>>()?;
1742 if key.is_empty() {
1743 return Err(Error::Source(
1744 "SQL Server incremental snapshot table has no primary-key fields".into(),
1745 ));
1746 }
1747 Ok(key)
1748}
1749
1750fn sqlserver_key_from_data_value(value: &DataValue) -> Result<SqlServerKeyValue> {
1751 match value {
1752 DataValue::Boolean(value) => Ok(SqlServerKeyValue::Boolean(*value)),
1753 DataValue::Int32(value) => Ok(SqlServerKeyValue::Int32(*value)),
1754 DataValue::Int64(value) => Ok(SqlServerKeyValue::Int64(*value)),
1755 DataValue::UInt64(value) => Ok(SqlServerKeyValue::UInt64(*value)),
1756 DataValue::Float64(value) if value.is_finite() => {
1757 Ok(SqlServerKeyValue::Float64(value.to_bits()))
1758 }
1759 DataValue::Decimal(value) => Ok(SqlServerKeyValue::Decimal(value.clone())),
1760 DataValue::String(value) => Ok(SqlServerKeyValue::String(value.clone())),
1761 DataValue::Bytes(value) => Ok(SqlServerKeyValue::Bytes(value.clone())),
1762 DataValue::Date(value) => Ok(SqlServerKeyValue::Date(value.clone())),
1763 DataValue::Time(value) => Ok(SqlServerKeyValue::Time(value.clone())),
1764 DataValue::Timestamp(value) => Ok(SqlServerKeyValue::Timestamp(value.clone())),
1765 DataValue::Uuid(value) => Ok(SqlServerKeyValue::Uuid(*value)),
1766 DataValue::Null => Err(Error::Source(
1767 "SQL Server incremental snapshot primary key contains NULL".into(),
1768 )),
1769 other => Err(Error::Source(format!(
1770 "SQL Server incremental snapshot primary key has unsupported value {other:?}"
1771 ))),
1772 }
1773}
1774
1775fn sqlserver_key_literal(value: &SqlServerKeyValue) -> Result<String> {
1776 match value {
1777 SqlServerKeyValue::Boolean(value) => Ok(u8::from(*value).to_string()),
1778 SqlServerKeyValue::Int32(value) => Ok(value.to_string()),
1779 SqlServerKeyValue::Int64(value) => Ok(value.to_string()),
1780 SqlServerKeyValue::UInt64(value) => Ok(value.to_string()),
1781 SqlServerKeyValue::Float64(bits) => {
1782 let value = f64::from_bits(*bits);
1783 if value.is_finite() {
1784 Ok(value.to_string())
1785 } else {
1786 Err(Error::State(
1787 "SQL Server incremental snapshot key contains a non-finite float".into(),
1788 ))
1789 }
1790 }
1791 SqlServerKeyValue::Decimal(value) => {
1792 let unsigned = value
1793 .strip_prefix('-')
1794 .or_else(|| value.strip_prefix('+'))
1795 .unwrap_or(value);
1796 let mut parts = unsigned.split('.');
1797 let whole = parts.next().unwrap_or_default();
1798 let fraction = parts.next();
1799 if whole.is_empty()
1800 || !whole.chars().all(|character| character.is_ascii_digit())
1801 || fraction.is_some_and(|fraction| {
1802 fraction.is_empty()
1803 || !fraction.chars().all(|character| character.is_ascii_digit())
1804 })
1805 || parts.next().is_some()
1806 {
1807 return Err(Error::State(format!(
1808 "SQL Server incremental snapshot has invalid decimal key {value:?}"
1809 )));
1810 }
1811 Ok(value.clone())
1812 }
1813 SqlServerKeyValue::Bytes(value) => Ok(format!("0x{}", hex::encode(value))),
1814 SqlServerKeyValue::String(value)
1815 | SqlServerKeyValue::Date(value)
1816 | SqlServerKeyValue::Time(value)
1817 | SqlServerKeyValue::Timestamp(value) => Ok(format!("N'{}'", quote_literal(value))),
1818 SqlServerKeyValue::Uuid(value) => Ok(format!("N'{value}'")),
1819 }
1820}
1821
1822fn sqlserver_key_predicate(
1823 columns: &[String],
1824 key: &[SqlServerKeyValue],
1825 greater: bool,
1826) -> Result<String> {
1827 validate_key_width(columns, key, if greater { "last" } else { "maximum" })?;
1828 let literals = key
1829 .iter()
1830 .map(sqlserver_key_literal)
1831 .collect::<Result<Vec<_>>>()?;
1832 let mut terms = Vec::new();
1833 for index in 0..columns.len() {
1834 let mut comparisons = (0..index)
1835 .map(|prefix| format!("{} = {}", columns[prefix], literals[prefix]))
1836 .collect::<Vec<_>>();
1837 comparisons.push(format!(
1838 "{} {} {}",
1839 columns[index],
1840 if greater { ">" } else { "<" },
1841 literals[index]
1842 ));
1843 terms.push(format!("({})", comparisons.join(" AND ")));
1844 }
1845 if !greater {
1846 terms.push(format!(
1847 "({})",
1848 columns
1849 .iter()
1850 .zip(&literals)
1851 .map(|(column, literal)| format!("{column} = {literal}"))
1852 .collect::<Vec<_>>()
1853 .join(" AND ")
1854 ));
1855 }
1856 Ok(format!("({})", terms.join(" OR ")))
1857}
1858
1859fn sqlserver_incremental_position(
1860 base: &SourcePosition,
1861 event_serial: u64,
1862) -> Result<SourcePosition> {
1863 let SourcePosition::SqlServer(position) = base else {
1864 return Err(Error::State(
1865 "SQL Server incremental snapshot requires a SQL Server position".into(),
1866 ));
1867 };
1868 Ok(SourcePosition::SqlServer(SqlServerPosition {
1869 database: position.database.clone(),
1870 commit_lsn: position.commit_lsn.clone(),
1871 change_lsn: position.change_lsn.clone(),
1872 event_serial: position.event_serial.saturating_add(event_serial),
1873 snapshot: false,
1874 }))
1875}
1876
1877async fn emit_incremental_checkpoint(
1878 output: &tokio::sync::mpsc::Sender<Result<SourceRecord>>,
1879 progress: Option<&IncrementalSnapshotProgress>,
1880 completed_signal_ids: &[String],
1881 position: &SourcePosition,
1882 acknowledgement: Option<SignalAcknowledgement>,
1883) -> Result<()> {
1884 output
1885 .send(Ok(SourceRecord {
1886 event: None,
1887 position: position.clone(),
1888 boundary: RecordBoundary::Heartbeat,
1889 connector_state: Some(encode_connector_state(progress, completed_signal_ids)?),
1890 signal_acknowledgements: acknowledgement.into_iter().collect(),
1891 }))
1892 .await
1893 .map_err(|_| Error::Cancelled)
1894}
1895
1896async fn execute_heartbeat_action(mut connection: SqlClient, query: String) -> Result<SqlClient> {
1897 connection
1898 .simple_query(&query)
1899 .await
1900 .map_err(sqlserver_error)?
1901 .into_results()
1902 .await
1903 .map_err(sqlserver_error)?;
1904 Ok(connection)
1905}
1906
1907fn sqlserver_heartbeat_record(
1908 connector_name: &str,
1909 database: &str,
1910 position: SourcePosition,
1911) -> SourceRecord {
1912 let observed_time = Utc::now();
1913 let mut attributes = BTreeMap::new();
1914 attributes.insert("rustium.heartbeat".into(), true.into());
1915 let mut after = Row::new();
1916 after.insert(
1917 "ts_ms".into(),
1918 DataValue::Int64(observed_time.timestamp_millis()),
1919 );
1920 let event = ChangeEvent {
1921 id: EventId::deterministic(
1922 connector_name,
1923 database,
1924 &position,
1925 "__heartbeat",
1926 u64::try_from(observed_time.timestamp_micros()).unwrap_or_default(),
1927 ),
1928 source: SourceMetadata {
1929 connector: "sqlserver".into(),
1930 connector_name: connector_name.into(),
1931 database: database.into(),
1932 schema: None,
1933 table: None,
1934 snapshot: false,
1935 version: CONNECTOR_VERSION.into(),
1936 attributes,
1937 },
1938 position: position.clone(),
1939 transaction: None,
1940 operation: Operation::Message,
1941 before: None,
1942 after: Some(after),
1943 schema: EventSchema {
1944 name: format!("{connector_name}.Heartbeat"),
1945 version: 1,
1946 fields: vec![FieldSchema {
1947 name: "ts_ms".into(),
1948 type_name: "int64".into(),
1949 optional: false,
1950 primary_key: false,
1951 }],
1952 },
1953 source_time: None,
1954 observed_time,
1955 };
1956 SourceRecord {
1957 event: Some(event),
1958 position,
1959 boundary: RecordBoundary::Heartbeat,
1960 connector_state: None,
1961 signal_acknowledgements: Vec::new(),
1962 }
1963}
1964
1965struct StreamingState {
1966 cursor: CdcCursor,
1967 transaction: Option<ActiveTransaction>,
1968 pending_updates: HashMap<(Vec<u8>, Vec<u8>, String), Row>,
1969 resume_position: Option<SourcePosition>,
1970}
1971
1972impl StreamingState {
1973 fn new(cursor: CdcCursor, resume_position: Option<SourcePosition>) -> Self {
1974 Self {
1975 cursor,
1976 transaction: None,
1977 pending_updates: HashMap::new(),
1978 resume_position,
1979 }
1980 }
1981
1982 fn transaction_metadata(
1983 &mut self,
1984 change: &RawChange,
1985 table: &CaptureTable,
1986 ) -> TransactionMetadata {
1987 if self
1988 .transaction
1989 .as_ref()
1990 .is_none_or(|transaction| transaction.commit_lsn != change.commit_lsn)
1991 {
1992 self.transaction = Some(ActiveTransaction {
1993 commit_lsn: change.commit_lsn.clone(),
1994 source_time: change.source_time,
1995 total_order: 0,
1996 collection_order: HashMap::new(),
1997 });
1998 }
1999 let transaction = self.transaction.as_mut().expect("transaction exists");
2000 transaction.total_order += 1;
2001 let collection_order = transaction.collection_order.entry(table.key()).or_insert(0);
2002 *collection_order += 1;
2003 TransactionMetadata {
2004 id: format_lsn(&change.commit_lsn),
2005 total_order: Some(transaction.total_order),
2006 collection_order: Some(*collection_order),
2007 }
2008 }
2009
2010 fn commit_record(&mut self, database: &str, commit_lsn: &[u8]) -> Option<SourceRecord> {
2011 self.transaction = None;
2012 let record = SourceRecord {
2013 event: None,
2014 position: sqlserver_position(
2015 database,
2016 commit_lsn,
2017 &max_lsn_bytes(),
2018 COMMIT_SERIAL,
2019 false,
2020 ),
2021 boundary: RecordBoundary::TransactionCommit,
2022 connector_state: None,
2023 signal_acknowledgements: Vec::new(),
2024 };
2025 if self.should_skip(&record.position) {
2026 None
2027 } else {
2028 Some(record)
2029 }
2030 }
2031
2032 fn should_skip(&mut self, position: &SourcePosition) -> bool {
2033 let Some(resume) = &self.resume_position else {
2034 return false;
2035 };
2036 if position.is_at_or_before(resume) {
2037 true
2038 } else {
2039 self.resume_position = None;
2040 false
2041 }
2042 }
2043
2044 fn safe_position(&self, database: &str) -> SourcePosition {
2045 sqlserver_position(
2046 database,
2047 &self.cursor.commit_lsn,
2048 &self.cursor.change_lsn,
2049 COMMIT_SERIAL,
2050 false,
2051 )
2052 }
2053}
2054
2055#[allow(clippy::too_many_arguments)]
2056async fn read_change_batch(
2057 client: &mut SqlClient,
2058 database: &str,
2059 connector_name: &str,
2060 captures: &[CaptureTable],
2061 state: &mut StreamingState,
2062 max_lsn: &[u8],
2063 fetch_size: usize,
2064 column_transformer: &ColumnTransformer,
2065 signal_table: Option<&(String, String)>,
2066) -> Result<Vec<SourceRecord>> {
2067 let query = change_query(captures, fetch_size.saturating_add(1));
2068 let rows = client
2069 .query(
2070 query,
2071 &[
2072 &state.cursor.commit_lsn,
2073 &state.cursor.change_lsn,
2074 &state.cursor.raw_operation,
2075 &max_lsn.to_vec(),
2076 ],
2077 )
2078 .await
2079 .map_err(sqlserver_error)?
2080 .into_first_result()
2081 .await
2082 .map_err(sqlserver_error)?;
2083
2084 if rows.is_empty() {
2085 if state.pending_updates.is_empty() {
2086 state.cursor = CdcCursor::at_snapshot(max_lsn.to_vec());
2087 return Ok(vec![SourceRecord {
2088 event: None,
2089 position: sqlserver_position(
2090 database,
2091 max_lsn,
2092 &max_lsn_bytes(),
2093 COMMIT_SERIAL,
2094 false,
2095 ),
2096 boundary: RecordBoundary::Heartbeat,
2097 connector_state: None,
2098 signal_acknowledgements: Vec::new(),
2099 }]);
2100 }
2101 return Ok(Vec::new());
2102 }
2103
2104 let has_lookahead = rows.len() > fetch_size;
2105 let process_count = rows.len().min(fetch_size);
2106 let mut raw_changes = Vec::with_capacity(process_count);
2107 for row in rows.iter().take(process_count) {
2108 raw_changes.push(raw_change(row, captures)?);
2109 }
2110 let lookahead_commit = has_lookahead
2111 .then(|| required_binary(&rows[process_count], "commit_lsn"))
2112 .transpose()?;
2113 let complete_last_commit = lookahead_commit.as_ref().is_none_or(|lookahead| {
2114 raw_changes
2115 .last()
2116 .is_some_and(|last| last.commit_lsn != *lookahead)
2117 });
2118 let capture_map = captures
2119 .iter()
2120 .map(|capture| (capture.capture_instance.as_str(), capture))
2121 .collect::<HashMap<_, _>>();
2122 let mut records = Vec::new();
2123
2124 for (index, change) in raw_changes.iter().enumerate() {
2125 let table = capture_map
2126 .get(change.capture_instance.as_str())
2127 .copied()
2128 .ok_or_else(|| {
2129 Error::Source(format!(
2130 "unknown SQL Server capture instance {:?}",
2131 change.capture_instance
2132 ))
2133 })?;
2134 let update_key = (
2135 change.commit_lsn.clone(),
2136 change.change_lsn.clone(),
2137 change.capture_instance.clone(),
2138 );
2139 let (operation, before, after, event_serial) = match change.raw_operation {
2140 RAW_DELETE => (Operation::Delete, Some(change.row.clone()), None, 1),
2141 RAW_INSERT => (Operation::Create, None, Some(change.row.clone()), 2),
2142 RAW_UPDATE_BEFORE => {
2143 state.pending_updates.insert(update_key, change.row.clone());
2144 state.cursor = CdcCursor {
2145 commit_lsn: change.commit_lsn.clone(),
2146 change_lsn: change.change_lsn.clone(),
2147 raw_operation: change.raw_operation,
2148 };
2149 continue;
2150 }
2151 RAW_UPDATE_AFTER => {
2152 let before = state.pending_updates.remove(&update_key).ok_or_else(|| {
2153 Error::Source(format!(
2154 "SQL Server update after-image at {}:{} has no before-image",
2155 format_lsn(&change.commit_lsn),
2156 format_lsn(&change.change_lsn)
2157 ))
2158 })?;
2159 (Operation::Update, Some(before), Some(change.row.clone()), 4)
2160 }
2161 operation => {
2162 return Err(Error::Source(format!(
2163 "unsupported SQL Server CDC operation {operation}"
2164 )));
2165 }
2166 };
2167 let position = sqlserver_position(
2168 database,
2169 &change.commit_lsn,
2170 &change.change_lsn,
2171 event_serial,
2172 false,
2173 );
2174 let transaction = state.transaction_metadata(change, table);
2175 let source_time = state
2176 .transaction
2177 .as_ref()
2178 .and_then(|transaction| transaction.source_time)
2179 .or(change.source_time);
2180 let mut event = ChangeEvent {
2181 id: EventId::deterministic(
2182 connector_name,
2183 database,
2184 &position,
2185 &format!("{database}.{}.{}", table.schema, table.table),
2186 event_serial,
2187 ),
2188 source: source_metadata(connector_name, database, table, false),
2189 position,
2190 transaction: Some(transaction),
2191 operation,
2192 before,
2193 after,
2194 schema: table.event_schema.clone(),
2195 source_time,
2196 observed_time: Utc::now(),
2197 };
2198 if signal_table != Some(&table.key()) {
2199 transform_event(column_transformer, &mut event, database, table);
2200 }
2201 let record = SourceRecord::data(event);
2202 if !state.should_skip(&record.position) {
2203 records.push(record);
2204 }
2205 state.cursor = CdcCursor {
2206 commit_lsn: change.commit_lsn.clone(),
2207 change_lsn: change.change_lsn.clone(),
2208 raw_operation: change.raw_operation,
2209 };
2210
2211 let next_commit = raw_changes.get(index + 1).map(|next| &next.commit_lsn);
2212 let commit_complete = next_commit.is_some_and(|next| next != &change.commit_lsn)
2213 || (index + 1 == raw_changes.len() && complete_last_commit);
2214 if commit_complete {
2215 if let Some(record) = state.commit_record(database, &change.commit_lsn) {
2216 records.push(record);
2217 }
2218 state.cursor = CdcCursor::at_snapshot(change.commit_lsn.clone());
2219 }
2220 }
2221 Ok(records)
2222}
2223
2224async fn connect(config: &SqlServerSourceConfig, database: &str) -> Result<SqlClient> {
2225 let mut tds = TdsConfig::new();
2226 tds.host(&config.hostname);
2227 tds.port(config.port);
2228 tds.database(database);
2229 tds.application_name("rustium");
2230 tds.authentication(AuthMethod::sql_server(&config.username, &config.password));
2231 if config.encrypt {
2232 tds.encryption(EncryptionLevel::Required);
2233 } else {
2234 tds.encryption(EncryptionLevel::NotSupported);
2235 }
2236 if config.trust_server_certificate {
2237 tds.trust_cert();
2238 }
2239 let address = tds.get_addr();
2240 let tcp = tokio::time::timeout(config.connect_timeout, TcpStream::connect(&address))
2241 .await
2242 .map_err(|_| Error::RetryableSource("timed out connecting to SQL Server".into()))?
2243 .map_err(sqlserver_io_error)?;
2244 tcp.set_nodelay(true).map_err(sqlserver_io_error)?;
2245 tokio::time::timeout(
2246 config.connect_timeout,
2247 Client::connect(tds, tcp.compat_write()),
2248 )
2249 .await
2250 .map_err(|_| Error::RetryableSource("timed out negotiating SQL Server TDS".into()))?
2251 .map_err(sqlserver_error)
2252}
2253
2254async fn discover_captures(
2255 client: &mut SqlClient,
2256 config: &SqlServerSourceConfig,
2257 connector_name: &str,
2258) -> Result<Vec<CaptureTable>> {
2259 let rows = client
2260 .simple_query(
2261 "SELECT s.name AS schema_name, t.name AS table_name, ct.capture_instance, \
2262 CAST(ct.source_object_id AS int) AS source_object_id \
2263 FROM cdc.change_tables ct \
2264 JOIN sys.tables t ON t.object_id = ct.source_object_id \
2265 JOIN sys.schemas s ON s.schema_id = t.schema_id \
2266 ORDER BY s.name, t.name, ct.capture_instance",
2267 )
2268 .await
2269 .map_err(sqlserver_error)?
2270 .into_first_result()
2271 .await
2272 .map_err(sqlserver_error)?;
2273 let mut captures = Vec::new();
2274 let mut selected_tables = HashMap::new();
2275 let signal_table = signal_table_key(config);
2276 for row in rows {
2277 let schema = required_string(&row, "schema_name")?;
2278 let table = required_string(&row, "table_name")?;
2279 let is_signal = signal_table.as_ref() == Some(&(schema.clone(), table.clone()));
2280 if !config.tables.includes(&schema, &table) && !is_signal {
2281 continue;
2282 }
2283 if selected_tables
2284 .insert((schema.clone(), table.clone()), ())
2285 .is_some()
2286 {
2287 return Err(Error::Configuration(format!(
2288 "SQL Server table {schema}.{table} has multiple CDC capture instances; Rustium requires an explicit single active instance"
2289 )));
2290 }
2291 let source_object_id = required::<i32>(&row, "source_object_id")?;
2292 let fields = discover_fields(client, source_object_id).await?;
2293 captures.push(CaptureTable {
2294 schema: schema.clone(),
2295 table: table.clone(),
2296 capture_instance: required_string(&row, "capture_instance")?,
2297 source_object_id,
2298 event_schema: EventSchema {
2299 name: format!("{connector_name}.{schema}.{table}.Envelope"),
2300 version: 1,
2301 fields,
2302 },
2303 });
2304 }
2305 Ok(captures)
2306}
2307
2308async fn discover_fields(client: &mut SqlClient, object_id: i32) -> Result<Vec<FieldSchema>> {
2309 let rows = client
2310 .query(
2311 "SELECT c.name AS column_name, ty.name AS type_name, CAST(c.max_length AS int) AS max_length, \
2312 CAST(c.precision AS int) AS precision, CAST(c.scale AS int) AS scale, \
2313 CAST(c.is_nullable AS bit) AS is_nullable, \
2314 CAST(CASE WHEN pk.column_id IS NULL THEN 0 ELSE 1 END AS bit) AS is_primary_key \
2315 FROM sys.columns c \
2316 JOIN sys.types ty ON ty.user_type_id = c.user_type_id \
2317 LEFT JOIN ( \
2318 SELECT ic.object_id, ic.column_id FROM sys.indexes i \
2319 JOIN sys.index_columns ic ON ic.object_id = i.object_id AND ic.index_id = i.index_id \
2320 WHERE i.is_primary_key = 1 \
2321 ) pk ON pk.object_id = c.object_id AND pk.column_id = c.column_id \
2322 WHERE c.object_id = @P1 ORDER BY c.column_id",
2323 &[&object_id],
2324 )
2325 .await
2326 .map_err(sqlserver_error)?
2327 .into_first_result()
2328 .await
2329 .map_err(sqlserver_error)?;
2330 let mut fields = Vec::new();
2331 for row in rows {
2332 let base = required_string(&row, "type_name")?;
2333 fields.push(FieldSchema {
2334 name: required_string(&row, "column_name")?,
2335 type_name: format_sql_type(
2336 &base,
2337 required::<i32>(&row, "max_length")?,
2338 required::<i32>(&row, "precision")?,
2339 required::<i32>(&row, "scale")?,
2340 ),
2341 optional: required::<bool>(&row, "is_nullable")?,
2342 primary_key: required::<bool>(&row, "is_primary_key")?,
2343 });
2344 }
2345 if fields.is_empty() {
2346 return Err(Error::Source(format!(
2347 "could not discover SQL Server columns for object_id {object_id}"
2348 )));
2349 }
2350 Ok(fields)
2351}
2352
2353async fn current_max_lsn(client: &mut SqlClient) -> Result<Vec<u8>> {
2354 let row = client
2355 .simple_query("SELECT sys.fn_cdc_get_max_lsn() AS max_lsn")
2356 .await
2357 .map_err(sqlserver_error)?
2358 .into_row()
2359 .await
2360 .map_err(sqlserver_error)?
2361 .ok_or_else(|| Error::Source("SQL Server returned no maximum CDC LSN".into()))?;
2362 required_binary(&row, "max_lsn")
2363}
2364
2365async fn validate_retention(
2366 client: &mut SqlClient,
2367 captures: &[CaptureTable],
2368 restart_lsn: &[u8],
2369) -> Result<()> {
2370 for capture in captures {
2371 let row = client
2372 .query(
2373 "SELECT sys.fn_cdc_get_min_lsn(@P1) AS min_lsn",
2374 &[&capture.capture_instance],
2375 )
2376 .await
2377 .map_err(sqlserver_error)?
2378 .into_row()
2379 .await
2380 .map_err(sqlserver_error)?
2381 .ok_or_else(|| Error::Source("SQL Server returned no minimum CDC LSN".into()))?;
2382 let min_lsn = required_binary(&row, "min_lsn")?;
2383 if restart_lsn != zero_lsn_bytes() && restart_lsn < min_lsn.as_slice() {
2384 return Err(Error::State(format!(
2385 "SQL Server CDC cleanup advanced capture instance {:?} to {}; checkpoint {} is no longer available",
2386 capture.capture_instance,
2387 format_lsn(&min_lsn),
2388 format_lsn(restart_lsn)
2389 )));
2390 }
2391 }
2392 Ok(())
2393}
2394
2395#[allow(clippy::too_many_arguments)]
2396async fn snapshot_table(
2397 client: &mut SqlClient,
2398 database: &str,
2399 connector_name: &str,
2400 capture: &CaptureTable,
2401 anchor: &[u8],
2402 ordinal: &mut u64,
2403 output: &tokio::sync::mpsc::Sender<Result<SourceRecord>>,
2404 column_transformer: &ColumnTransformer,
2405) -> Result<()> {
2406 let values = capture
2407 .event_schema
2408 .fields
2409 .iter()
2410 .enumerate()
2411 .map(|(index, field)| change_value_expression(index, field))
2412 .collect::<Vec<_>>()
2413 .join(", ");
2414 let primary_key = capture
2415 .event_schema
2416 .fields
2417 .iter()
2418 .filter(|field| field.primary_key)
2419 .map(|field| format!("ct.{}", quote_identifier(&field.name)))
2420 .collect::<Vec<_>>();
2421 let ordering = if primary_key.is_empty() {
2422 String::new()
2423 } else {
2424 format!(" ORDER BY {}", primary_key.join(", "))
2425 };
2426 let query = format!(
2427 "SELECT {values} FROM {}.{} AS ct{ordering}",
2428 quote_identifier(&capture.schema),
2429 quote_identifier(&capture.table)
2430 );
2431 let mut rows = client
2432 .simple_query(query)
2433 .await
2434 .map_err(sqlserver_error)?
2435 .into_row_stream();
2436 let mut count = 0_u64;
2437 while let Some(row) = rows.try_next().await.map_err(sqlserver_error)? {
2438 *ordinal += 1;
2439 count += 1;
2440 let position = sqlserver_position(database, anchor, &zero_lsn_bytes(), *ordinal, true);
2441 let mut event = ChangeEvent {
2442 id: EventId::deterministic(
2443 connector_name,
2444 database,
2445 &position,
2446 &format!("{database}.{}.{}", capture.schema, capture.table),
2447 *ordinal,
2448 ),
2449 source: source_metadata(connector_name, database, capture, true),
2450 position,
2451 transaction: None,
2452 operation: Operation::Read,
2453 before: None,
2454 after: Some(convert_tds_row(&row, &capture.event_schema)?),
2455 schema: capture.event_schema.clone(),
2456 source_time: None,
2457 observed_time: Utc::now(),
2458 };
2459 transform_event(column_transformer, &mut event, database, capture);
2460 output
2461 .send(Ok(SourceRecord::data(event)))
2462 .await
2463 .map_err(|_| Error::Cancelled)?;
2464 }
2465 drop(rows);
2466 info!(table = %format!("{}.{}", capture.schema, capture.table), rows = count, "SQL Server snapshot table completed");
2467 Ok(())
2468}
2469
2470fn change_query(captures: &[CaptureTable], limit: usize) -> String {
2471 let unions = captures
2472 .iter()
2473 .map(|capture| {
2474 let values = capture
2475 .event_schema
2476 .fields
2477 .iter()
2478 .enumerate()
2479 .map(|(index, field)| change_value_expression(index, field))
2480 .collect::<Vec<_>>()
2481 .join(", ");
2482 format!(
2483 "SELECT ct.__$start_lsn AS commit_lsn, ct.__$seqval AS change_lsn, \
2484 CAST(ct.__$operation AS int) AS operation, \
2485 CAST(N'{}' AS nvarchar(128)) AS capture_instance, \
2486 (SELECT {values} FOR JSON PATH, WITHOUT_ARRAY_WRAPPER, INCLUDE_NULL_VALUES) AS row_data, \
2487 sys.fn_cdc_map_lsn_to_time(ct.__$start_lsn) AS source_time \
2488 FROM cdc.{} ct",
2489 quote_literal(&capture.capture_instance),
2490 quote_identifier(&format!("{}_CT", capture.capture_instance)),
2491 )
2492 })
2493 .collect::<Vec<_>>()
2494 .join(" UNION ALL ");
2495 format!(
2496 "SELECT TOP ({limit}) commit_lsn, change_lsn, operation, capture_instance, row_data, source_time \
2497 FROM ({unions}) AS changes \
2498 WHERE (commit_lsn > @P1 OR (commit_lsn = @P1 AND \
2499 (change_lsn > @P2 OR (change_lsn = @P2 AND operation > @P3)))) \
2500 AND commit_lsn <= @P4 \
2501 ORDER BY commit_lsn, change_lsn, operation"
2502 )
2503}
2504
2505fn change_value_expression(index: usize, field: &FieldSchema) -> String {
2506 let column = format!("ct.{}", quote_identifier(&field.name));
2507 let base = base_type(&field.type_name);
2508 let value = match base {
2509 "binary" | "varbinary" | "image" | "rowversion" | "timestamp" => {
2510 format!("CONVERT(varchar(max), {column}, 2)")
2511 }
2512 "date" | "time" | "datetime" | "datetime2" | "smalldatetime" | "datetimeoffset" => {
2513 format!("CONVERT(nvarchar(64), {column}, 126)")
2514 }
2515 "geometry" | "geography" => {
2516 format!("CONVERT(varchar(max), {column}.Serialize(), 2)")
2517 }
2518 "hierarchyid" => format!("{column}.ToString()"),
2519 _ => format!("CONVERT(nvarchar(max), {column})"),
2520 };
2521 format!("{value} AS [c{index}]")
2522}
2523
2524fn raw_change(row: &TdsRow, captures: &[CaptureTable]) -> Result<RawChange> {
2525 let capture_instance = required_string(row, "capture_instance")?;
2526 let capture = captures
2527 .iter()
2528 .find(|capture| capture.capture_instance == capture_instance)
2529 .ok_or_else(|| {
2530 Error::Source(format!(
2531 "SQL Server returned unknown capture instance {capture_instance:?}"
2532 ))
2533 })?;
2534 let json = required_string(row, "row_data")?;
2535 Ok(RawChange {
2536 commit_lsn: required_binary(row, "commit_lsn")?,
2537 change_lsn: required_binary(row, "change_lsn")?,
2538 raw_operation: required::<i32>(row, "operation")?,
2539 capture_instance,
2540 row: convert_json_row(&json, &capture.event_schema)?,
2541 source_time: row
2542 .try_get::<NaiveDateTime, _>("source_time")
2543 .map_err(sqlserver_error)?
2544 .map(|time| DateTime::from_naive_utc_and_offset(time, Utc)),
2545 })
2546}
2547
2548fn convert_json_row(json: &str, schema: &EventSchema) -> Result<Row> {
2549 let value: serde_json::Value = serde_json::from_str(json)?;
2550 let object = value
2551 .as_object()
2552 .ok_or_else(|| Error::Source("SQL Server CDC row JSON is not an object".into()))?;
2553 Ok(schema
2554 .fields
2555 .iter()
2556 .enumerate()
2557 .map(|(index, field)| {
2558 let value = object
2559 .get(&format!("c{index}"))
2560 .map_or(DataValue::Null, |value| {
2561 json_text_value(value, &field.type_name)
2562 });
2563 (field.name.clone(), value)
2564 })
2565 .collect())
2566}
2567
2568fn json_text_value(value: &serde_json::Value, type_name: &str) -> DataValue {
2569 match value {
2570 serde_json::Value::Null => DataValue::Null,
2571 serde_json::Value::String(value) => convert_sqlserver_text(value, type_name),
2572 serde_json::Value::Bool(value) => DataValue::Boolean(*value),
2573 serde_json::Value::Number(value) => DataValue::Decimal(value.to_string()),
2574 _ => DataValue::Json(value.clone()),
2575 }
2576}
2577
2578fn convert_sqlserver_text(value: &str, type_name: &str) -> DataValue {
2579 match base_type(type_name) {
2580 "bit" => DataValue::Boolean(matches!(value, "1" | "true" | "TRUE")),
2581 "tinyint" | "smallint" | "int" => value
2582 .parse::<i32>()
2583 .map_or_else(|_| DataValue::String(value.into()), DataValue::Int32),
2584 "bigint" => value
2585 .parse::<i64>()
2586 .map_or_else(|_| DataValue::String(value.into()), DataValue::Int64),
2587 "real" | "float" => value
2588 .parse::<f64>()
2589 .map_or_else(|_| DataValue::String(value.into()), DataValue::Float64),
2590 "decimal" | "numeric" | "money" | "smallmoney" => DataValue::Decimal(value.into()),
2591 "binary" | "varbinary" | "image" | "rowversion" | "timestamp" | "geometry"
2592 | "geography" => {
2593 hex::decode(value).map_or_else(|_| DataValue::String(value.into()), DataValue::Bytes)
2594 }
2595 "date" => DataValue::Date(value.into()),
2596 "time" => DataValue::Time(value.into()),
2597 "datetime" | "datetime2" | "smalldatetime" | "datetimeoffset" => {
2598 DataValue::Timestamp(value.into())
2599 }
2600 "uniqueidentifier" => uuid::Uuid::parse_str(value)
2601 .map_or_else(|_| DataValue::String(value.into()), DataValue::Uuid),
2602 _ => DataValue::String(value.into()),
2603 }
2604}
2605
2606fn convert_tds_row(row: &TdsRow, schema: &EventSchema) -> Result<Row> {
2607 if row.len() != schema.fields.len() {
2608 return Err(Error::Source(format!(
2609 "SQL Server snapshot row has {} values but schema has {} fields",
2610 row.len(),
2611 schema.fields.len()
2612 )));
2613 }
2614 schema
2615 .fields
2616 .iter()
2617 .enumerate()
2618 .map(|(index, field)| {
2619 Ok((
2620 field.name.clone(),
2621 convert_tds_value(row, index, &field.type_name)?,
2622 ))
2623 })
2624 .collect()
2625}
2626
2627fn convert_tds_value(row: &TdsRow, index: usize, type_name: &str) -> Result<DataValue> {
2628 let cell = row
2629 .cells()
2630 .nth(index)
2631 .map(|(_, value)| value)
2632 .ok_or_else(|| Error::Source("SQL Server row cell is missing".into()))?;
2633 let value = match cell {
2634 ColumnData::U8(value) => {
2635 value.map_or(DataValue::Null, |value| DataValue::Int32(value.into()))
2636 }
2637 ColumnData::I16(value) => {
2638 value.map_or(DataValue::Null, |value| DataValue::Int32(value.into()))
2639 }
2640 ColumnData::I32(value) => value.map_or(DataValue::Null, DataValue::Int32),
2641 ColumnData::I64(value) => value.map_or(DataValue::Null, DataValue::Int64),
2642 ColumnData::F32(value) => {
2643 value.map_or(DataValue::Null, |value| DataValue::Float64(value.into()))
2644 }
2645 ColumnData::F64(value) => value.map_or(DataValue::Null, DataValue::Float64),
2646 ColumnData::Bit(value) => value.map_or(DataValue::Null, DataValue::Boolean),
2647 ColumnData::String(value) => value.as_ref().map_or(DataValue::Null, |value| {
2648 convert_sqlserver_text(value, type_name)
2649 }),
2650 ColumnData::Guid(value) => value.map_or(DataValue::Null, DataValue::Uuid),
2651 ColumnData::Binary(value) => value
2652 .as_ref()
2653 .map_or(DataValue::Null, |value| DataValue::Bytes(value.to_vec())),
2654 ColumnData::Numeric(value) => value.map_or(DataValue::Null, |value| {
2655 DataValue::Decimal(value.to_string())
2656 }),
2657 ColumnData::Xml(value) => value.as_ref().map_or(DataValue::Null, |value| {
2658 DataValue::String(value.to_string())
2659 }),
2660 ColumnData::DateTime(None)
2661 | ColumnData::SmallDateTime(None)
2662 | ColumnData::Time(None)
2663 | ColumnData::Date(None)
2664 | ColumnData::DateTime2(None)
2665 | ColumnData::DateTimeOffset(None) => DataValue::Null,
2666 ColumnData::DateTime(Some(_))
2667 | ColumnData::SmallDateTime(Some(_))
2668 | ColumnData::DateTime2(Some(_)) => row
2669 .try_get::<NaiveDateTime, _>(index)
2670 .map_err(sqlserver_error)?
2671 .map_or(DataValue::Null, |value| {
2672 DataValue::Timestamp(sqlserver_datetime(value))
2673 }),
2674 ColumnData::Time(Some(_)) => row
2675 .try_get::<NaiveTime, _>(index)
2676 .map_err(sqlserver_error)?
2677 .map_or(DataValue::Null, |value| DataValue::Time(value.to_string())),
2678 ColumnData::Date(Some(_)) => row
2679 .try_get::<NaiveDate, _>(index)
2680 .map_err(sqlserver_error)?
2681 .map_or(DataValue::Null, |value| DataValue::Date(value.to_string())),
2682 ColumnData::DateTimeOffset(Some(_)) => row
2683 .try_get::<DateTime<FixedOffset>, _>(index)
2684 .map_err(sqlserver_error)?
2685 .map_or(DataValue::Null, |value| {
2686 DataValue::Timestamp(value.to_rfc3339())
2687 }),
2688 };
2689 Ok(value)
2690}
2691
2692fn sqlserver_datetime(value: NaiveDateTime) -> String {
2693 value.to_string().replacen(' ', "T", 1)
2694}
2695
2696fn source_metadata(
2697 connector_name: &str,
2698 database: &str,
2699 capture: &CaptureTable,
2700 snapshot: bool,
2701) -> SourceMetadata {
2702 let mut attributes = BTreeMap::new();
2703 attributes.insert(
2704 "capture_instance".into(),
2705 capture.capture_instance.clone().into(),
2706 );
2707 attributes.insert("source_object_id".into(), capture.source_object_id.into());
2708 SourceMetadata {
2709 connector: "sqlserver".into(),
2710 connector_name: connector_name.into(),
2711 database: database.into(),
2712 schema: Some(capture.schema.clone()),
2713 table: Some(capture.table.clone()),
2714 snapshot,
2715 version: CONNECTOR_VERSION.into(),
2716 attributes,
2717 }
2718}
2719
2720fn sqlserver_position(
2721 database: &str,
2722 commit_lsn: &[u8],
2723 change_lsn: &[u8],
2724 event_serial: u64,
2725 snapshot: bool,
2726) -> SourcePosition {
2727 SourcePosition::SqlServer(SqlServerPosition {
2728 database: database.into(),
2729 commit_lsn: format_lsn(commit_lsn),
2730 change_lsn: format_lsn(change_lsn),
2731 event_serial,
2732 snapshot,
2733 })
2734}
2735
2736fn snapshot_begin_sql(mode: &str) -> &'static str {
2737 match mode {
2738 "exclusive" => "SET TRANSACTION ISOLATION LEVEL SERIALIZABLE; BEGIN TRANSACTION",
2739 "snapshot" => "SET TRANSACTION ISOLATION LEVEL SNAPSHOT; BEGIN TRANSACTION",
2740 "read_committed" => "SET TRANSACTION ISOLATION LEVEL READ COMMITTED; BEGIN TRANSACTION",
2741 "read_uncommitted" => "SET TRANSACTION ISOLATION LEVEL READ UNCOMMITTED; BEGIN TRANSACTION",
2742 _ => "SET TRANSACTION ISOLATION LEVEL REPEATABLE READ; BEGIN TRANSACTION",
2743 }
2744}
2745
2746fn format_sql_type(base: &str, max_length: i32, precision: i32, scale: i32) -> String {
2747 match base {
2748 "decimal" | "numeric" => format!("{base}({precision},{scale})"),
2749 "char" | "varchar" | "binary" | "varbinary" => {
2750 if max_length < 0 {
2751 format!("{base}(max)")
2752 } else {
2753 format!("{base}({max_length})")
2754 }
2755 }
2756 "nchar" | "nvarchar" => {
2757 if max_length < 0 {
2758 format!("{base}(max)")
2759 } else {
2760 format!("{base}({})", max_length / 2)
2761 }
2762 }
2763 "datetime2" | "datetimeoffset" | "time" => format!("{base}({scale})"),
2764 _ => base.into(),
2765 }
2766}
2767
2768fn required<'a, T>(row: &'a TdsRow, name: &str) -> Result<T>
2769where
2770 T: tiberius::FromSql<'a>,
2771{
2772 row.try_get(name)
2773 .map_err(sqlserver_error)?
2774 .ok_or_else(|| Error::Source(format!("SQL Server result {name:?} is null")))
2775}
2776
2777fn required_string(row: &TdsRow, name: &str) -> Result<String> {
2778 required::<&str>(row, name).map(str::to_string)
2779}
2780
2781fn required_binary(row: &TdsRow, name: &str) -> Result<Vec<u8>> {
2782 let value = required::<&[u8]>(row, name)?.to_vec();
2783 if value.len() != LSN_SIZE {
2784 return Err(Error::Source(format!(
2785 "SQL Server {name} has {} bytes, expected {LSN_SIZE}",
2786 value.len()
2787 )));
2788 }
2789 Ok(value)
2790}
2791
2792fn quote_identifier(identifier: &str) -> String {
2793 format!("[{}]", identifier.replace(']', "]]"))
2794}
2795
2796fn quote_literal(value: &str) -> String {
2797 value.replace('\'', "''")
2798}
2799
2800fn base_type(type_name: &str) -> &str {
2801 type_name.split('(').next().unwrap_or(type_name)
2802}
2803
2804fn format_lsn(lsn: &[u8]) -> String {
2805 format!("0x{}", hex::encode_upper(lsn))
2806}
2807
2808fn parse_lsn(lsn: &str) -> Result<Vec<u8>> {
2809 let bytes = hex::decode(lsn.strip_prefix("0x").unwrap_or(lsn)).map_err(|error| {
2810 Error::State(format!(
2811 "invalid SQL Server checkpoint LSN {lsn:?}: {error}"
2812 ))
2813 })?;
2814 if bytes.len() != LSN_SIZE {
2815 return Err(Error::State(format!(
2816 "invalid SQL Server checkpoint LSN length {}; expected {LSN_SIZE}",
2817 bytes.len()
2818 )));
2819 }
2820 Ok(bytes)
2821}
2822
2823fn zero_lsn_bytes() -> Vec<u8> {
2824 vec![0; LSN_SIZE]
2825}
2826
2827fn max_lsn_bytes() -> Vec<u8> {
2828 vec![u8::MAX; LSN_SIZE]
2829}
2830
2831fn sqlserver_error(error: tiberius::error::Error) -> Error {
2832 let retryable = matches!(error, tiberius::error::Error::Io { .. })
2833 || error.code().is_some_and(is_transient_sqlserver_code);
2834 if retryable {
2835 Error::RetryableSource(error.to_string())
2836 } else {
2837 Error::Source(error.to_string())
2838 }
2839}
2840
2841fn sqlserver_io_error(error: std::io::Error) -> Error {
2842 Error::RetryableSource(error.to_string())
2843}
2844
2845const fn is_transient_sqlserver_code(code: u32) -> bool {
2846 matches!(
2847 code,
2848 233 | 1205
2849 | 1222
2850 | 10_053
2851 | 10_054
2852 | 10_060
2853 | 10_928
2854 | 10_929
2855 | 40_197
2856 | 40_501
2857 | 40_613
2858 | 49_918
2859 | 49_919
2860 | 49_920
2861 )
2862}
2863
2864#[cfg(test)]
2865mod tests {
2866 use super::*;
2867
2868 #[test]
2869 fn qualifies_sqlserver_snapshot_collection_filters() {
2870 let snapshot = SnapshotConfig {
2871 include_collections: vec![r"inventory\.dbo\.orders".into()],
2872 ..SnapshotConfig::default()
2873 };
2874 assert!(snapshot_includes(&snapshot, "inventory", "dbo", "orders"));
2875 assert!(!snapshot_includes(&snapshot, "archive", "dbo", "orders"));
2876 assert!(!snapshot_includes(
2877 &snapshot,
2878 "inventory",
2879 "dbo",
2880 "orders_history"
2881 ));
2882 }
2883
2884 #[test]
2885 fn round_trips_sqlserver_lsn() {
2886 let bytes = vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 255];
2887 assert_eq!(parse_lsn(&format_lsn(&bytes)).unwrap(), bytes);
2888 }
2889
2890 #[test]
2891 fn distinguishes_complete_and_mid_transaction_cursors() {
2892 let complete = CdcCursor::at_snapshot(vec![1; LSN_SIZE]);
2893 assert!(complete.commit_complete());
2894
2895 let incomplete = CdcCursor {
2896 commit_lsn: vec![1; LSN_SIZE],
2897 change_lsn: vec![2; LSN_SIZE],
2898 raw_operation: RAW_UPDATE_BEFORE,
2899 };
2900 assert!(!incomplete.commit_complete());
2901 }
2902
2903 #[test]
2904 fn classifies_only_transient_sqlserver_failures_for_retry() {
2905 let io_error = tiberius::error::Error::Io {
2906 kind: std::io::ErrorKind::ConnectionReset,
2907 message: "connection reset".into(),
2908 };
2909 assert!(matches!(
2910 sqlserver_error(io_error),
2911 Error::RetryableSource(message) if message.contains("connection reset")
2912 ));
2913 assert!(is_transient_sqlserver_code(1205));
2914 assert!(is_transient_sqlserver_code(40_613));
2915 assert!(!is_transient_sqlserver_code(208));
2916
2917 assert!(matches!(
2918 sqlserver_error(tiberius::error::Error::Protocol("invalid token".into())),
2919 Error::Source(message) if message.contains("invalid token")
2920 ));
2921 }
2922
2923 #[test]
2924 fn enforces_sqlserver_retry_policy_boundaries() {
2925 let disabled = RetryPolicy {
2926 max_retries: 0,
2927 ..RetryPolicy::default()
2928 };
2929 assert!(!sqlserver_retry_allowed(&disabled, 0));
2930
2931 let finite = RetryPolicy {
2932 max_retries: 2,
2933 ..RetryPolicy::default()
2934 };
2935 assert!(sqlserver_retry_allowed(&finite, 0));
2936 assert!(sqlserver_retry_allowed(&finite, 1));
2937 assert!(!sqlserver_retry_allowed(&finite, 2));
2938
2939 let unbounded = RetryPolicy {
2940 max_retries: -1,
2941 ..RetryPolicy::default()
2942 };
2943 assert!(sqlserver_retry_allowed(&unbounded, u64::MAX));
2944 }
2945
2946 #[test]
2947 fn converts_sqlserver_values() {
2948 assert_eq!(
2949 convert_sqlserver_text("12.30", "decimal(10,2)"),
2950 DataValue::Decimal("12.30".into())
2951 );
2952 assert_eq!(convert_sqlserver_text("1", "bit"), DataValue::Boolean(true));
2953 assert_eq!(
2954 convert_sqlserver_text("00FF", "varbinary(2)"),
2955 DataValue::Bytes(vec![0, 255])
2956 );
2957 assert_eq!(
2958 convert_sqlserver_text("E6100000010C0000000000000040000000000000F03F", "geometry"),
2959 DataValue::Bytes(hex::decode("E6100000010C0000000000000040000000000000F03F").unwrap())
2960 );
2961 assert_eq!(
2962 change_value_expression(
2963 0,
2964 &FieldSchema {
2965 name: "location".into(),
2966 type_name: "geography".into(),
2967 optional: false,
2968 primary_key: false,
2969 }
2970 ),
2971 "CONVERT(varchar(max), ct.[location].Serialize(), 2) AS [c0]"
2972 );
2973 assert_eq!(
2974 sqlserver_datetime(
2975 NaiveDate::from_ymd_opt(2026, 7, 16)
2976 .unwrap()
2977 .and_hms_micro_opt(9, 30, 45, 123_400)
2978 .unwrap()
2979 ),
2980 "2026-07-16T09:30:45.123400"
2981 );
2982 }
2983
2984 #[test]
2985 fn builds_bounded_direct_cdc_query() {
2986 let capture = CaptureTable {
2987 schema: "dbo".into(),
2988 table: "orders".into(),
2989 capture_instance: "dbo_orders".into(),
2990 source_object_id: 42,
2991 event_schema: EventSchema {
2992 name: "orders".into(),
2993 version: 1,
2994 fields: vec![FieldSchema {
2995 name: "id".into(),
2996 type_name: "bigint".into(),
2997 optional: false,
2998 primary_key: true,
2999 }],
3000 },
3001 };
3002 let query = change_query(&[capture], 513);
3003 assert!(query.contains("TOP (513)"));
3004 assert!(query.contains("cdc.[dbo_orders_CT]"));
3005 assert!(query.contains("ORDER BY commit_lsn, change_lsn, operation"));
3006 }
3007
3008 #[test]
3009 fn builds_composite_sqlserver_keyset_predicates() {
3010 let columns = vec!["ct.[tenant_id]".into(), "ct.[id]".into()];
3011 let key = vec![SqlServerKeyValue::Int32(7), SqlServerKeyValue::Int64(42)];
3012 assert_eq!(
3013 sqlserver_key_predicate(&columns, &key, true).unwrap(),
3014 "((ct.[tenant_id] > 7) OR (ct.[tenant_id] = 7 AND ct.[id] > 42))"
3015 );
3016 assert_eq!(
3017 sqlserver_key_predicate(&columns, &key, false).unwrap(),
3018 "((ct.[tenant_id] < 7) OR (ct.[tenant_id] = 7 AND ct.[id] < 42) OR (ct.[tenant_id] = 7 AND ct.[id] = 42))"
3019 );
3020 }
3021
3022 #[test]
3023 fn validates_sqlserver_signal_table_layout() {
3024 let mut capture = CaptureTable {
3025 schema: "dbo".into(),
3026 table: "rustium_signal".into(),
3027 capture_instance: "dbo_rustium_signal".into(),
3028 source_object_id: 42,
3029 event_schema: EventSchema {
3030 name: "signal".into(),
3031 version: 1,
3032 fields: ["id", "type", "data"]
3033 .into_iter()
3034 .map(|name| FieldSchema {
3035 name: name.into(),
3036 type_name: "nvarchar(max)".into(),
3037 optional: false,
3038 primary_key: name == "id",
3039 })
3040 .collect(),
3041 },
3042 };
3043 assert!(validate_signal_schema(&capture).is_ok());
3044 capture.event_schema.fields[0].type_name = "nvarchar(20)".into();
3045 assert!(matches!(
3046 validate_signal_schema(&capture),
3047 Err(Error::Configuration(message)) if message.contains("at least 42 characters")
3048 ));
3049 }
3050
3051 #[test]
3052 fn builds_sqlserver_heartbeat_record_at_commit_position() {
3053 let position = sqlserver_position(
3054 "inventory",
3055 &[1; LSN_SIZE],
3056 &max_lsn_bytes(),
3057 COMMIT_SERIAL,
3058 false,
3059 );
3060 let record =
3061 sqlserver_heartbeat_record("inventory-sqlserver", "inventory", position.clone());
3062 let event = record.event.unwrap();
3063 assert_eq!(record.boundary, RecordBoundary::Heartbeat);
3064 assert_eq!(record.position, position);
3065 assert_eq!(event.operation, Operation::Message);
3066 assert_eq!(event.source.table, None);
3067 assert_eq!(
3068 event.source.attributes.get("rustium.heartbeat"),
3069 Some(&true.into())
3070 );
3071 assert!(matches!(
3072 event.after.unwrap().get("ts_ms"),
3073 Some(DataValue::Int64(_))
3074 ));
3075 }
3076}