1use std::collections::{BTreeMap, BTreeSet};
13use std::rc::Rc;
14use std::sync::Arc;
15use std::time::Instant;
16
17use differential_dataflow::AsCollection;
18use futures::StreamExt;
19use itertools::Itertools;
20use mz_ore::cast::CastFrom;
21use mz_ore::collections::HashMap;
22use mz_ore::future::InTask;
23use mz_repr::{Diff, GlobalId, Row, RowArena};
24use mz_sql_server_util::SqlServerCdcMetrics;
25use mz_sql_server_util::cdc::{CdcEvent, Lsn, Operation as CdcOperation};
26use mz_sql_server_util::desc::SqlServerRowDecoder;
27use mz_sql_server_util::inspect::{
28 ensure_database_cdc_enabled, ensure_sql_server_agent_running, get_latest_restore_history_id,
29};
30use mz_storage_types::dyncfgs::SQL_SERVER_SOURCE_VALIDATE_RESTORE_HISTORY;
31use mz_storage_types::errors::{DataflowError, DecodeError, DecodeErrorKind};
32use mz_storage_types::sources::SqlServerSourceConnection;
33use mz_storage_types::sources::sql_server::{MAX_LSN_WAIT, SNAPSHOT_PROGRESS_REPORT_INTERVAL};
34use mz_timely_util::builder_async::{
35 AsyncOutputHandle, OperatorBuilder as AsyncOperatorBuilder, PressOnDropButton,
36};
37use mz_timely_util::containers::stack::FueledBuilder;
38use timely::container::CapacityContainerBuilder;
39use timely::dataflow::operators::vec::Map;
40use timely::dataflow::operators::{CapabilitySet, Concat};
41use timely::dataflow::{Scope, StreamVec};
42use timely::progress::{Antichain, Timestamp};
43
44use crate::metrics::source::sql_server::SqlServerSourceMetrics;
45use crate::source::RawSourceCreationConfig;
46use crate::source::sql_server::{
47 DefiniteError, ReplicationError, SourceOutputInfo, TransientError,
48};
49use crate::source::types::{FuelSize, SignaledFuture, SourceMessage, StackedCollection};
50
51static REPL_READER: &str = "reader";
57
58pub(crate) fn render<'scope>(
59 scope: Scope<'scope, Lsn>,
60 config: RawSourceCreationConfig,
61 outputs: BTreeMap<GlobalId, SourceOutputInfo>,
62 source: SqlServerSourceConnection,
63 metrics: SqlServerSourceMetrics,
64) -> (
65 StackedCollection<'scope, Lsn, (u64, Result<SourceMessage, DataflowError>)>,
66 StreamVec<'scope, Lsn, ReplicationError>,
67 PressOnDropButton,
68) {
69 let op_name = format!("SqlServerReplicationReader({})", config.id);
70 let mut builder = AsyncOperatorBuilder::new(op_name, scope);
71
72 let (data_output, data_stream) = builder.new_output::<FueledBuilder<_>>();
73
74 let (definite_error_handle, definite_errors) =
76 builder.new_output::<CapacityContainerBuilder<_>>();
77
78 let (button, transient_errors) = builder.build_fallible(move |caps| {
79 let busy_signal = Arc::clone(&config.busy_signal);
80 Box::pin(SignaledFuture::new(busy_signal, async move {
81 let [
82 data_cap_set,
83 definite_error_cap_set,
84 ]: &mut [_; 2] = caps.try_into().unwrap();
85
86 let connection_config = source
87 .connection
88 .resolve_config(
89 &config.config.connection_context.secrets_reader,
90 &config.config,
91 InTask::Yes,
92 )
93 .await?;
94 let mut client = mz_sql_server_util::Client::connect(connection_config).await?;
95
96 let worker_id = config.worker_id;
97
98 let mut decoder_map: BTreeMap<_, _> = BTreeMap::new();
100 let mut capture_instance_to_snapshot: BTreeMap<Arc<str>, Vec<_>> = BTreeMap::new();
102 let mut capture_instances: BTreeMap<Arc<str>, Vec<_>> = BTreeMap::new();
104 let mut export_statistics: BTreeMap<_, Vec<_>> = BTreeMap::new();
106 let mut included_columns: HashMap<u64, Vec<Arc<str>>> = HashMap::new();
109
110 for (export_id, output) in outputs.iter() {
111 let key = output.partition_index;
112 if decoder_map.insert(key, Arc::clone(&output.decoder)).is_some() {
113 panic!("Multiple decoders for output index {}", output.partition_index);
114 }
115 let included_cols = output.decoder.included_column_names();
121 included_columns.insert(output.partition_index, included_cols);
122
123 capture_instances
124 .entry(Arc::clone(&output.capture_instance))
125 .or_default()
126 .push(output.partition_index);
127
128 if *output.resume_upper == [Lsn::minimum()] {
129 capture_instance_to_snapshot
130 .entry(Arc::clone(&output.capture_instance))
131 .or_default()
132 .push((output.partition_index, output.initial_lsn));
133 }
134 export_statistics.entry(Arc::clone(&output.capture_instance))
135 .or_default()
136 .push(
137 config
138 .statistics
139 .get(export_id)
140 .expect("statistics have been intialized")
141 .clone(),
142 );
143 }
144
145 metrics.snapshot_table_count.set(u64::cast_from(capture_instance_to_snapshot.len()));
150 if !capture_instance_to_snapshot.is_empty() {
151 for stats in config.statistics.values() {
152 stats.set_snapshot_records_known(0);
153 stats.set_snapshot_records_staged(0);
154 }
155 }
156 if !config.responsible_for(REPL_READER) {
159 return Ok::<_, TransientError>(());
160 }
161
162 let snapshot_instances = capture_instance_to_snapshot
163 .keys()
164 .map(|i| i.as_ref());
165
166 let snapshot_tables =
168 mz_sql_server_util::inspect::get_tables_for_capture_instance(
169 &mut client,
170 snapshot_instances,
171 )
172 .await?;
173
174 let current_restore_history_id = get_latest_restore_history_id(&mut client).await?;
176 if current_restore_history_id != source.extras.restore_history_id {
177 if SQL_SERVER_SOURCE_VALIDATE_RESTORE_HISTORY.get(config.config.config_set()) {
178 let definite_error = DefiniteError::RestoreHistoryChanged(
179 source.extras.restore_history_id.clone(),
180 current_restore_history_id.clone()
181 );
182 tracing::warn!(?definite_error, "Restore detected, exiting");
183
184 return_definite_error(
185 definite_error,
186 capture_instances.values().flat_map(|indexes| indexes.iter().copied()),
187 data_output,
188 data_cap_set,
189 definite_error_handle,
190 definite_error_cap_set,
191 ).await;
192 return Ok(());
193 } else {
194 tracing::warn!(
195 "Restore history mismatch ignored: expected={expected:?} actual={actual:?}",
196 expected=source.extras.restore_history_id,
197 actual=current_restore_history_id
198 );
199 }
200 }
201
202 ensure_database_cdc_enabled(&mut client).await?;
205 ensure_sql_server_agent_running(&mut client).await?;
206
207 for table in &snapshot_tables {
211 let qualified_table_name = format!("{schema_name}.{table_name}",
212 schema_name = table.schema_name,
213 table_name = table.name);
214 let size_calc_start = Instant::now();
215 let table_total =
216 mz_sql_server_util::inspect::snapshot_size(
217 &mut client,
218 &table.schema_name,
219 &table.name,
220 )
221 .await?;
222 metrics.set_snapshot_table_size_latency(
223 &qualified_table_name,
224 size_calc_start.elapsed().as_secs_f64()
225 );
226 for export_stat in export_statistics.get(&table.capture_instance.name).unwrap() {
227 export_stat.set_snapshot_records_known(u64::cast_from(table_total));
228 export_stat.set_snapshot_records_staged(0);
229 }
230 }
231 let cdc_metrics = PrometheusSqlServerCdcMetrics{inner: &metrics};
232 let mut cdc_handle = client
233 .cdc(capture_instances.keys().cloned(), cdc_metrics)
234 .max_lsn_wait(MAX_LSN_WAIT.get(config.config.config_set()));
235
236 let snapshot_lsns: BTreeMap<Arc<str>, Lsn> = {
239 cdc_handle.wait_for_ready().await?;
242
243 tracing::info!(%config.worker_id, "timely-{worker_id} upstream is ready");
247
248 let report_interval =
249 SNAPSHOT_PROGRESS_REPORT_INTERVAL.handle(config.config.config_set());
250 let mut last_report = Instant::now();
251 let mut snapshot_lsns = BTreeMap::new();
252
253 for table in snapshot_tables {
254 let (snapshot_lsn, snapshot) = cdc_handle
256 .snapshot(&table, config.worker_id, config.id)
257 .await?;
258
259 tracing::info!(
260 %config.id,
261 %table.name,
262 %table.schema_name,
263 %snapshot_lsn,
264 "timely-{worker_id} snapshot start",
265 );
266
267 let mut snapshot = std::pin::pin!(snapshot);
268
269 snapshot_lsns.insert(
270 Arc::clone(&table.capture_instance.name),
271 snapshot_lsn,
272 );
273
274 let ci_name = &table.capture_instance.name;
275 let partition_indexes = capture_instance_to_snapshot
276 .get(ci_name)
277 .unwrap_or_else(|| {
278 panic!(
279 "no snapshot outputs in known capture \
280 instances [{}] for capture instance: \
281 '{}'",
282 capture_instance_to_snapshot
283 .keys()
284 .join(","),
285 ci_name,
286 );
287 });
288
289 let mut snapshot_staged = 0;
290 while let Some(result) = snapshot.next().await {
291 let sql_server_row =
292 result.map_err(TransientError::from)?;
293
294 if last_report.elapsed() > report_interval.get() {
295 last_report = Instant::now();
296 let stats =
297 export_statistics.get(ci_name).unwrap();
298 for export_stat in stats {
299 export_stat.set_snapshot_records_staged(
300 snapshot_staged,
301 );
302 }
303 }
304
305 for (partition_idx, _) in partition_indexes {
306 let mut mz_row = Row::default();
308 let arena = RowArena::default();
309
310 let decoder = decoder_map
311 .get(partition_idx)
312 .expect("decoder for output");
313 let message = decode(
316 decoder,
317 &sql_server_row,
318 &mut mz_row,
319 &arena,
320 None,
321 );
322 let update =
323 ((*partition_idx, message), Lsn::minimum(), Diff::ONE);
324 let size = update.fuel_size();
325 data_output
326 .give_fueled(&data_cap_set[0], update, size)
327 .await;
328 }
329 snapshot_staged += 1;
330 }
331
332 tracing::info!(
333 %config.id,
334 %table.name,
335 %table.schema_name,
336 %snapshot_lsn,
337 "timely-{worker_id} snapshot complete",
338 );
339 metrics.snapshot_table_count.dec();
340 let stats = export_statistics.get(ci_name).unwrap();
343 for export_stat in stats {
344 export_stat.set_snapshot_records_staged(snapshot_staged);
345 export_stat.set_snapshot_records_known(snapshot_staged);
346 }
347 }
348
349 snapshot_lsns
350 };
351
352 let mut rewinds: BTreeMap<_, _> = capture_instance_to_snapshot
367 .iter()
368 .flat_map(|(capture_instance, export_ids)|{
369 let snapshot_lsn = snapshot_lsns.get(capture_instance).expect("snapshot lsn must be collected for capture instance");
370 export_ids
371 .iter()
372 .map(|(idx, initial_lsn)| (*idx, (*initial_lsn, *snapshot_lsn)))
373 }).collect();
374
375 for (initial_lsn, snapshot_lsn) in rewinds.values() {
381 assert!(
382 initial_lsn <= snapshot_lsn,
383 "initial_lsn={initial_lsn} snapshot_lsn={snapshot_lsn}"
384 );
385 }
386
387 tracing::debug!("rewinds to process: {rewinds:?}");
388
389 capture_instance_to_snapshot.clear();
390
391 let mut resume_lsns = BTreeMap::new();
393 for src_info in outputs.values() {
394 let resume_lsn = src_info.resume_lsn_or(src_info.initial_lsn.increment());
398 resume_lsns.entry(Arc::clone(&src_info.capture_instance))
399 .and_modify(|existing| *existing = std::cmp::min(*existing, resume_lsn))
400 .or_insert(resume_lsn);
401 }
402
403 tracing::info!(%config.id, ?resume_lsns, "timely-{} replication starting", config.worker_id);
404 for instance in capture_instances.keys() {
405 let resume_lsn = resume_lsns
406 .get(instance)
407 .expect("resume_lsn exists for capture instance");
408 cdc_handle = cdc_handle.start_lsn(instance, *resume_lsn);
409 }
410
411 let cdc_stream = cdc_handle
413 .poll_interval(config.timestamp_interval)
414 .into_stream();
415 let mut cdc_stream = std::pin::pin!(cdc_stream);
416
417 let mut errored_partitions = BTreeSet::new();
418
419 let mut log_rewinds_complete = true;
423
424 let mut deferred_updates = BTreeMap::new();
439
440 while let Some(event) = cdc_stream.next().await {
441 let event = event.map_err(TransientError::from)?;
442 tracing::trace!(?config.id, ?event, "got replication event");
443
444 tracing::trace!("deferred_updates = {deferred_updates:?}");
445 match event {
446 CdcEvent::Progress { next_lsn } => {
449 tracing::debug!(?config.id, ?next_lsn, "got a closed lsn");
450 rewinds.retain(|_, (_, snapshot_lsn)| next_lsn <= *snapshot_lsn);
453 if rewinds.is_empty() {
454 if log_rewinds_complete {
455 tracing::debug!("rewinds complete");
456 log_rewinds_complete = false;
457 }
458 data_cap_set.downgrade(Antichain::from_elem(next_lsn));
459 } else {
460 tracing::debug!("rewinds remaining: {:?}", rewinds);
461 }
462
463 if let Some(((deferred_lsn, _seqval), _row)) =
466 deferred_updates.first_key_value()
467 && *deferred_lsn < next_lsn
468 {
469 panic!(
470 "deferred update lsn {deferred_lsn} \
471 < progress lsn {next_lsn}: {:?}",
472 deferred_updates.keys()
473 );
474 }
475
476 }
477 CdcEvent::Data {
479 capture_instance,
480 lsn,
481 changes,
482 } => {
483 let Some(partition_indexes) =
484 capture_instances.get(&capture_instance)
485 else {
486 let definite_error =
487 DefiniteError::ProgrammingError(format!(
488 "capture instance didn't exist: \
489 '{capture_instance}'"
490 ));
491 return_definite_error(
492 definite_error,
493 capture_instances
494 .values()
495 .flat_map(|indexes| {
496 indexes.iter().copied()
497 }),
498 data_output,
499 data_cap_set,
500 definite_error_handle,
501 definite_error_cap_set,
502 )
503 .await;
504 return Ok(());
505 };
506
507 let (valid_partitions, err_partitions) =
508 partition_indexes
509 .iter()
510 .partition::<Vec<u64>, _>(
511 |&partition_idx| {
512 !errored_partitions
513 .contains(partition_idx)
514 },
515 );
516
517 if err_partitions.len() > 0 {
518 metrics.ignored.inc_by(u64::cast_from(changes.len()));
519 }
520
521 handle_data_event(
522 changes,
523 &valid_partitions,
524 &decoder_map,
525 lsn,
526 &rewinds,
527 &data_output,
528 data_cap_set,
529 &metrics,
530 &mut deferred_updates,
531 ).await?
532 },
533 CdcEvent::SchemaUpdate {
534 capture_instance,
535 table,
536 ddl_event,
537 } => {
538 let Some(partition_indexes) =
539 capture_instances.get(&capture_instance)
540 else {
541 let definite_error =
542 DefiniteError::ProgrammingError(format!(
543 "capture instance didn't exist: \
544 '{capture_instance}'"
545 ));
546 return_definite_error(
547 definite_error,
548 capture_instances
549 .values()
550 .flat_map(|indexes| {
551 indexes.iter().copied()
552 }),
553 data_output,
554 data_cap_set,
555 definite_error_handle,
556 definite_error_cap_set,
557 )
558 .await;
559 return Ok(());
560 };
561 let error =
562 DefiniteError::IncompatibleSchemaChange(
563 capture_instance.to_string(),
564 table.to_string(),
565 );
566 for partition_idx in partition_indexes {
567 let cols = included_columns
568 .get(partition_idx)
569 .unwrap_or_else(|| {
570 panic!(
571 "Partition index didn't \
572 exist: '{partition_idx}'"
573 )
574 });
575 if !errored_partitions
576 .contains(partition_idx)
577 && !ddl_event.is_compatible(cols)
578 {
579 let msg = Err(
580 error.clone().into(),
581 );
582 let update = (
583 (*partition_idx, msg),
584 ddl_event.lsn,
585 Diff::ONE,
586 );
587 let size = update.fuel_size();
588 data_output
589 .give_fueled(&data_cap_set[0], update, size)
590 .await;
591 errored_partitions.insert(*partition_idx);
592 }
593 }
594 }
595 };
596 }
597 Err(TransientError::ReplicationEOF)
598 }))
599 });
600
601 let error_stream = definite_errors.concat(transient_errors.map(ReplicationError::Transient));
602
603 (
604 data_stream.as_collection(),
605 error_stream,
606 button.press_on_drop(),
607 )
608}
609
610async fn handle_data_event(
611 changes: Vec<CdcOperation>,
612 partition_indexes: &[u64],
613 decoder_map: &BTreeMap<u64, Arc<SqlServerRowDecoder>>,
614 commit_lsn: Lsn,
615 rewinds: &BTreeMap<u64, (Lsn, Lsn)>,
616 data_output: &StackedAsyncOutputHandle<Lsn, (u64, Result<SourceMessage, DataflowError>)>,
617 data_cap_set: &CapabilitySet<Lsn>,
618 metrics: &SqlServerSourceMetrics,
619 deferred_updates: &mut BTreeMap<(Lsn, Lsn), CdcOperation>,
620) -> Result<(), TransientError> {
621 let mut mz_row = Row::default();
622 let arena = RowArena::default();
623
624 for change in changes {
625 let mut deferred_update: Option<_> = None;
629 let (sql_server_row, diff): (_, _) = match change {
630 CdcOperation::Insert(sql_server_row) => {
631 metrics.inserts.inc();
632 (sql_server_row, Diff::ONE)
633 }
634 CdcOperation::Delete(sql_server_row) => {
635 metrics.deletes.inc();
636 (sql_server_row, Diff::MINUS_ONE)
637 }
638
639 CdcOperation::UpdateNew(seqval, sql_server_row) => {
642 metrics.updates.inc();
644 deferred_update = deferred_updates.remove(&(commit_lsn, seqval));
645 if deferred_update.is_none() {
646 tracing::trace!("capture deferred UpdateNew ({commit_lsn}, {seqval})");
647 deferred_updates.insert(
648 (commit_lsn, seqval),
649 CdcOperation::UpdateNew(seqval, sql_server_row),
650 );
651 continue;
652 }
653 (sql_server_row, Diff::ZERO)
655 }
656 CdcOperation::UpdateOld(seqval, sql_server_row) => {
657 deferred_update = deferred_updates.remove(&(commit_lsn, seqval));
658 if deferred_update.is_none() {
659 tracing::trace!("capture deferred UpdateOld ({commit_lsn}, {seqval})");
660 deferred_updates.insert(
661 (commit_lsn, seqval),
662 CdcOperation::UpdateOld(seqval, sql_server_row),
663 );
664 continue;
665 }
666 (sql_server_row, Diff::ZERO)
668 }
669 };
670
671 for partition_idx in partition_indexes {
673 let decoder = decoder_map.get(partition_idx).unwrap();
674
675 let rewind = rewinds.get(partition_idx);
676 if rewind.is_some_and(|(initial_lsn, _)| commit_lsn <= *initial_lsn) {
679 continue;
680 }
681
682 let (message, diff) = if let Some(ref deferred_update) = deferred_update {
683 let (old_row, new_row) = match deferred_update {
684 CdcOperation::UpdateOld(_seqval, row) => (row, &sql_server_row),
685 CdcOperation::UpdateNew(_seqval, row) => (&sql_server_row, row),
686 CdcOperation::Insert(_) | CdcOperation::Delete(_) => unreachable!(),
687 };
688
689 let update_old = decode(decoder, old_row, &mut mz_row, &arena, Some(new_row));
690 if rewind.is_some_and(|(_, snapshot_lsn)| commit_lsn <= *snapshot_lsn) {
691 let update = (
692 (*partition_idx, update_old.clone()),
693 Lsn::minimum(),
694 Diff::ONE,
695 );
696 let size = update.fuel_size();
697 data_output
698 .give_fueled(&data_cap_set[0], update, size)
699 .await;
700 }
701 let update = ((*partition_idx, update_old), commit_lsn, Diff::MINUS_ONE);
702 let size = update.fuel_size();
703 data_output
704 .give_fueled(&data_cap_set[0], update, size)
705 .await;
706
707 (
708 decode(decoder, new_row, &mut mz_row, &arena, None),
709 Diff::ONE,
710 )
711 } else {
712 (
713 decode(decoder, &sql_server_row, &mut mz_row, &arena, None),
714 diff,
715 )
716 };
717 assert_ne!(Diff::ZERO, diff);
718 if rewind.is_some_and(|(_, snapshot_lsn)| commit_lsn <= *snapshot_lsn) {
719 let update = ((*partition_idx, message.clone()), Lsn::minimum(), -diff);
720 let size = update.fuel_size();
721 data_output
722 .give_fueled(&data_cap_set[0], update, size)
723 .await;
724 }
725 let update = ((*partition_idx, message), commit_lsn, diff);
726 let size = update.fuel_size();
727 data_output
728 .give_fueled(&data_cap_set[0], update, size)
729 .await;
730 }
731 }
732 Ok(())
733}
734
735type StackedAsyncOutputHandle<T, D> =
736 AsyncOutputHandle<T, FueledBuilder<CapacityContainerBuilder<Vec<(D, T, Diff)>>>>;
737
738fn decode(
741 decoder: &SqlServerRowDecoder,
742 row: &tiberius::Row,
743 mz_row: &mut Row,
744 arena: &RowArena,
745 new_row: Option<&tiberius::Row>,
746) -> Result<SourceMessage, DataflowError> {
747 match decoder.decode(row, mz_row, arena, new_row) {
748 Ok(()) => Ok(SourceMessage {
749 key: Row::default(),
750 value: mz_row.clone(),
751 metadata: Row::default(),
752 }),
753 Err(e) => {
754 let kind = DecodeErrorKind::Text(e.to_string().into());
755 let raw = format!("{row:?}");
757 Err(DataflowError::DecodeError(Box::new(DecodeError {
758 kind,
759 raw: raw.as_bytes().to_vec(),
760 })))
761 }
762 }
763}
764
765async fn return_definite_error(
767 err: DefiniteError,
768 outputs: impl Iterator<Item = u64>,
769 data_handle: StackedAsyncOutputHandle<Lsn, (u64, Result<SourceMessage, DataflowError>)>,
770 data_capset: &CapabilitySet<Lsn>,
771 errs_handle: AsyncOutputHandle<Lsn, CapacityContainerBuilder<Vec<ReplicationError>>>,
772 errs_capset: &CapabilitySet<Lsn>,
773) {
774 for output_idx in outputs {
775 let update = (
776 (output_idx, Err(err.clone().into())),
777 Lsn {
781 vlf_id: u32::MAX,
782 block_id: u32::MAX,
783 record_id: u16::MAX,
784 },
785 Diff::ONE,
786 );
787 let size = update.fuel_size();
788 data_handle.give_fueled(&data_capset[0], update, size).await;
789 }
790 errs_handle.give(
791 &errs_capset[0],
792 ReplicationError::DefiniteError(Rc::new(err)),
793 );
794}
795
796struct PrometheusSqlServerCdcMetrics<'a> {
798 inner: &'a SqlServerSourceMetrics,
799}
800
801impl<'a> SqlServerCdcMetrics for PrometheusSqlServerCdcMetrics<'a> {
802 fn snapshot_table_lock_start(&self, table_name: &str) {
803 self.inner.update_snapshot_table_lock_count(table_name, 1);
804 }
805
806 fn snapshot_table_lock_end(&self, table_name: &str) {
807 self.inner.update_snapshot_table_lock_count(table_name, -1);
808 }
809}