Skip to main content

rustium_sqlserver/
source.rs

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}