1use std::borrow::Cow;
4
5use arrow_array::RecordBatch;
6use futures_util::io::{AsyncRead, AsyncWrite};
7
8use crate::observability::{
9 DIRECT_ENCODING_PHASE, TARGET_METADATA_VALIDATION_PHASE, WRITER_INITIALIZATION_PHASE,
10 writer::{BatchWriteTrace, DirectRawBatchObserver, FinishTrace, WriterInitializationTrace},
11};
12use crate::{
13 Diagnostic, DiagnosticCode, DiagnosticSet, FieldRef, PlannedSchema, Result, SchemaMapping,
14 TableName, WritePhase,
15};
16
17use super::{
18 SchemaCheck,
19 context::RuntimeConversionContext,
20 direct::{
21 DirectEncoder, MeasuredDirectBatch, MeasuredRowRange,
22 plan::{DirectColumnEncoding, DirectColumnPlan, DirectEncoderPlan},
23 },
24 profile,
25 record_batch::RecordBatchView,
26 token_row::tiberius_row_owned,
27};
28use crate::conversion::arrow_to_mssql::{
29 fixed_size_binary::FixedSizeBinaryArrowToMssql, primitive::PrimitiveArrowToMssql,
30 temporal::TemporalArrowToMssql, variable_width::VariableWidthArrowToMssql,
31};
32
33const DIRECT_RAW_MAX_PAYLOAD_BYTES: usize = 8 * 1024 * 1024;
34
35#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
37pub enum WriteBackend {
38 #[default]
40 Auto,
41 BaselineTokenRow,
43 DirectFramedBulk,
45 DirectRawBulk,
47}
48
49#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
51pub struct WriteOptions {
52 pub backend: WriteBackend,
54 pub schema_check: SchemaCheck,
56}
57
58#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash)]
60pub struct WriteStats {
61 pub rows_written: u64,
63 pub batches_written: u64,
65}
66
67#[derive(Debug)]
68struct WriterState {
69 backend: WriteBackend,
70 direct_encoder: Option<DirectEncoder>,
71 schema_check: SchemaCheck,
72 runtime_context: RuntimeConversionContext,
73 mappings: Vec<SchemaMapping>,
74 stats: WriteStats,
75}
76
77impl WriterState {
78 fn new(
79 requested_backend: WriteBackend,
80 schema_check: SchemaCheck,
81 planned_schema: PlannedSchema,
82 ) -> Result<Self> {
83 let backend = resolve_backend(requested_backend)?;
84 let runtime_context =
85 RuntimeConversionContext::new(planned_schema.profile(), planned_schema.plan_options());
86 let mappings = planned_schema.into_mappings();
87 let direct_encoder = match backend {
88 WriteBackend::DirectFramedBulk | WriteBackend::DirectRawBulk => {
89 Some(DirectEncoder::new_with_context(&mappings, runtime_context)?)
90 }
91 WriteBackend::Auto | WriteBackend::BaselineTokenRow => None,
92 };
93
94 Ok(Self {
95 backend,
96 direct_encoder,
97 schema_check,
98 runtime_context,
99 mappings,
100 stats: WriteStats::default(),
101 })
102 }
103
104 fn backend(&self) -> WriteBackend {
105 self.backend
106 }
107
108 fn direct_encoder(&self) -> Option<&DirectEncoder> {
109 self.direct_encoder.as_ref()
110 }
111
112 fn mappings(&self) -> &[SchemaMapping] {
113 &self.mappings
114 }
115
116 fn schema_check(&self) -> SchemaCheck {
117 self.schema_check
118 }
119
120 fn runtime_context(&self) -> RuntimeConversionContext {
121 self.runtime_context
122 }
123
124 fn stats(&self) -> WriteStats {
125 self.stats
126 }
127
128 fn record_accepted_batch(&mut self, rows: u64) -> WriteStats {
129 self.stats.rows_written = self.stats.rows_written.saturating_add(rows);
130 self.stats.batches_written = self.stats.batches_written.saturating_add(1);
131 self.stats
132 }
133}
134
135#[derive(Debug)]
137pub struct BulkWriter<'client, S>
138where
139 S: AsyncRead + AsyncWrite + Unpin + Send,
140{
141 state: WriterState,
142 request: tiberius::BulkLoadRequest<'client, S>,
143}
144
145impl<'client, S> BulkWriter<'client, S>
146where
147 S: AsyncRead + AsyncWrite + Unpin + Send,
148{
149 pub async fn new(
151 client: &'client mut tiberius::Client<S>,
152 table: TableName,
153 planned_schema: PlannedSchema,
154 options: WriteOptions,
155 ) -> Result<Self> {
156 let mapping_count = planned_schema.mappings().len();
157 let mut trace = WriterInitializationTrace::new(&table, options.backend, mapping_count);
158 trace.emit_started();
159
160 let state = match WriterState::new(options.backend, options.schema_check, planned_schema) {
161 Ok(state) => state,
162 Err(err) => {
163 trace.emit_failed(WRITER_INITIALIZATION_PHASE, &err);
164 return Err(err.with_write_phase(WritePhase::WriterInitialization));
165 }
166 };
167 trace.record_resolved_backend(state.backend());
168 trace.record_direct_target_validation_required(matches!(
169 state.backend(),
170 WriteBackend::DirectFramedBulk | WriteBackend::DirectRawBulk
171 ));
172
173 let mut request = match state.backend() {
174 WriteBackend::BaselineTokenRow
175 | WriteBackend::DirectFramedBulk
176 | WriteBackend::DirectRawBulk => {
177 let table_sql = bulk_insert_table_sql(&table);
178 let columns = match client
179 .bulk_insert_columns(&table_sql)
180 .await
181 .map_err(|source| crate::Error::Tiberius { source })
182 {
183 Ok(columns) => columns,
184 Err(err) => {
185 trace.emit_failed(TARGET_METADATA_VALIDATION_PHASE, &err);
186 return Err(err.with_write_phase(WritePhase::TargetMetadataValidation));
187 }
188 };
189 trace.emit_target_metadata_validation_started();
190 if let Err(err) = validate_bulk_target_columns(columns.iter(), state.mappings()) {
191 trace.emit_failed(TARGET_METADATA_VALIDATION_PHASE, &err);
192 return Err(err.with_write_phase(WritePhase::TargetMetadataValidation));
193 }
194 if matches!(
195 state.backend(),
196 WriteBackend::DirectFramedBulk | WriteBackend::DirectRawBulk
197 ) {
198 let encoder = match state.direct_encoder().ok_or_else(|| {
199 crate::Error::BackendUnavailable {
200 backend: state.backend(),
201 reason: "direct bulk encoder is not available for this writer"
202 .to_owned(),
203 }
204 }) {
205 Ok(encoder) => encoder,
206 Err(err) => {
207 trace.emit_failed(TARGET_METADATA_VALIDATION_PHASE, &err);
208 return Err(err.with_write_phase(WritePhase::TargetMetadataValidation));
209 }
210 };
211 if let Err(err) =
212 validate_direct_bulk_target_column_types(columns.iter(), encoder.plan())
213 {
214 trace.emit_failed(TARGET_METADATA_VALIDATION_PHASE, &err);
215 return Err(err.with_write_phase(WritePhase::TargetMetadataValidation));
216 }
217 }
218 trace.emit_target_metadata_validation_completed();
219 match client
220 .bulk_insert_with_columns(&table_sql, columns)
221 .await
222 .map_err(|source| crate::Error::Tiberius { source })
223 {
224 Ok(request) => request,
225 Err(err) => {
226 trace.emit_failed(WRITER_INITIALIZATION_PHASE, &err);
227 return Err(err.with_write_phase(WritePhase::WriterInitialization));
228 }
229 }
230 }
231 WriteBackend::Auto => {
232 let err = execution_unavailable(state.backend());
233 trace.emit_failed(WRITER_INITIALIZATION_PHASE, &err);
234 return Err(err.with_write_phase(WritePhase::WriterInitialization));
235 }
236 };
237
238 if state.backend() == WriteBackend::DirectRawBulk {
239 request.enable_direct_packet_writes();
240 }
241
242 trace.emit_completed();
243
244 Ok(Self { state, request })
245 }
246
247 pub async fn write_batch(&mut self, batch: &RecordBatch) -> Result<WriteStats> {
249 match self.state.backend() {
250 WriteBackend::BaselineTokenRow => {
251 write_traced_batch_to_sink(&mut self.state, &mut self.request, batch).await
252 }
253 WriteBackend::DirectFramedBulk | WriteBackend::DirectRawBulk => {
254 write_traced_direct_batch_to_sink(&mut self.state, &mut self.request, batch).await
255 }
256 WriteBackend::Auto => Err(execution_unavailable(WriteBackend::Auto)),
257 }
258 }
259
260 pub async fn finish(self) -> Result<WriteStats> {
262 let Self { state, request } = self;
263 finish_writer_to_sink(state, request).await
264 }
265}
266
267async fn finish_writer_to_sink<Sink>(state: WriterState, sink: Sink) -> Result<WriteStats>
268where
269 Sink: FinishSink,
270{
271 let trace = FinishTrace::new(state.backend(), state.stats());
272 trace.emit_started();
273 let stats = state.stats();
274
275 if let Err(err) = sink.finalize_bulk_load().await {
276 trace.emit_failed(&err);
277 return Err(err.with_write_phase(WritePhase::Finalize));
278 }
279
280 trace.emit_completed();
281 Ok(stats)
282}
283
284trait FinishSink {
285 async fn finalize_bulk_load(self) -> Result<()>;
286}
287
288impl<S> FinishSink for tiberius::BulkLoadRequest<'_, S>
289where
290 S: AsyncRead + AsyncWrite + Unpin + Send,
291{
292 async fn finalize_bulk_load(self) -> Result<()> {
293 #[cfg(feature = "bench-profile")]
294 {
295 let (_result, stats) = self
296 .finalize_with_stats()
297 .await
298 .map_err(|source| crate::Error::Tiberius { source })?;
299 profile::record_bulk_load_stats(stats);
300 }
301
302 #[cfg(not(feature = "bench-profile"))]
303 self.finalize()
304 .await
305 .map_err(|source| crate::Error::Tiberius { source })?;
306
307 Ok(())
308 }
309}
310
311fn bulk_insert_table_sql(table: &TableName) -> String {
312 table.quoted_sql()
313}
314
315fn record_batch_view<'a>(
316 batch: &'a RecordBatch,
317 mappings: &'a [SchemaMapping],
318 schema_check: SchemaCheck,
319 runtime_context: RuntimeConversionContext,
320) -> Result<RecordBatchView<'a>> {
321 match schema_check {
322 SchemaCheck::Strict => RecordBatchView::new_with_context(batch, mappings, runtime_context),
323 }
324}
325
326fn validate_batch_rows(view: &RecordBatchView<'_>) -> Result<()> {
327 for row_index in 0..view.row_count() {
328 let _cells = view.mssql_row(row_index)?;
329 }
330
331 Ok(())
332}
333
334fn validate_bulk_target_columns<Column>(
335 columns: impl ExactSizeIterator<Item = Column>,
336 mappings: &[SchemaMapping],
337) -> Result<()>
338where
339 Column: BulkTargetColumnMetadata,
340{
341 let column_count = columns.len();
342 let mut diagnostics = DiagnosticSet::new();
343
344 if column_count != mappings.len() {
345 diagnostics.push(Diagnostic::error(
346 DiagnosticCode::SchemaMismatch,
347 format!(
348 "bulk target has {column_count} updateable column(s) but mappings contain {} column(s)",
349 mappings.len()
350 ),
351 ));
352 }
353
354 for (position, (column, mapping)) in columns.zip(mappings).enumerate() {
355 validate_bulk_target_column(position, column, mapping, &mut diagnostics);
356 }
357
358 if diagnostics.has_errors() {
359 return Err(crate::Error::ValueConversion { diagnostics });
360 }
361
362 Ok(())
363}
364
365fn validate_bulk_target_column(
366 position: usize,
367 column: impl BulkTargetColumnMetadata,
368 mapping: &SchemaMapping,
369 diagnostics: &mut DiagnosticSet,
370) {
371 if column.ordinal() != position {
372 diagnostics.push(bulk_target_column_diagnostic(
373 mapping,
374 format!(
375 "bulk target column ordinal {} does not match mapping position {position}",
376 column.ordinal()
377 ),
378 ));
379 }
380
381 if column.name() != mapping.mssql().name().as_str() {
382 diagnostics.push(bulk_target_column_diagnostic(
383 mapping,
384 format!(
385 "bulk target column name {} does not match planned MSSQL column name {}",
386 column.name(),
387 mapping.mssql().name().as_str()
388 ),
389 ));
390 }
391
392 if column.is_nullable() != mapping.mssql().nullable() {
393 diagnostics.push(bulk_target_column_diagnostic(
394 mapping,
395 format!(
396 "bulk target column nullability {} does not match planned MSSQL column nullability {}",
397 column.is_nullable(),
398 mapping.mssql().nullable()
399 ),
400 ));
401 }
402}
403
404fn validate_direct_bulk_target_column_types<Column>(
405 columns: impl ExactSizeIterator<Item = Column>,
406 plan: &DirectEncoderPlan,
407) -> Result<()>
408where
409 Column: BulkTargetColumnMetadata,
410{
411 let column_count = columns.len();
412 let mut diagnostics = DiagnosticSet::new();
413
414 if column_count != plan.column_count() {
415 diagnostics.push(Diagnostic::error(
416 DiagnosticCode::SchemaMismatch,
417 format!(
418 "bulk target has {column_count} updateable column(s) but direct plan contains {} column(s)",
419 plan.column_count()
420 ),
421 ));
422 }
423
424 for (column, plan_column) in columns.zip(plan.columns()) {
425 validate_direct_bulk_target_column_type(column, plan_column, &mut diagnostics);
426 }
427
428 if diagnostics.has_errors() {
429 return Err(crate::Error::ValueConversion { diagnostics });
430 }
431
432 Ok(())
433}
434
435fn validate_direct_bulk_target_column_type(
436 column: impl BulkTargetColumnMetadata,
437 plan_column: &DirectColumnPlan,
438 diagnostics: &mut DiagnosticSet,
439) {
440 let Some(expected) = expected_direct_bulk_column_type(plan_column) else {
441 diagnostics.push(
442 Diagnostic::error(
443 DiagnosticCode::DirectEncodingUnsupportedMapping,
444 format!(
445 "direct target type validation is not implemented for {:?}",
446 plan_column.encoding()
447 ),
448 )
449 .with_field(FieldRef::new(
450 plan_column.source_index(),
451 plan_column.source_name(),
452 )),
453 );
454 return;
455 };
456 let actual = column.column_type();
457
458 if actual != expected
459 && !matches!(
460 (actual, expected),
461 (
462 tiberius::ColumnType::Datetime,
463 tiberius::ColumnType::Datetimen
464 )
465 )
466 {
467 diagnostics.push(
468 Diagnostic::error(
469 DiagnosticCode::SchemaMismatch,
470 format!(
471 "bulk target column type {actual:?} does not match direct encoder type {expected:?}"
472 ),
473 )
474 .with_field(FieldRef::new(
475 plan_column.source_index(),
476 plan_column.source_name(),
477 )),
478 );
479 }
480
481 if let Some((expected_precision, expected_scale)) =
482 expected_direct_decimal_precision_scale(plan_column)
483 {
484 match column.decimal_precision_scale() {
485 Some((actual_precision, actual_scale))
486 if actual_precision == expected_precision && actual_scale == expected_scale => {}
487 Some((actual_precision, actual_scale)) => diagnostics.push(
488 Diagnostic::error(
489 DiagnosticCode::SchemaMismatch,
490 format!(
491 "bulk target decimal precision/scale ({actual_precision},{actual_scale}) does not match direct encoder precision/scale ({expected_precision},{expected_scale})"
492 ),
493 )
494 .with_field(FieldRef::new(
495 plan_column.source_index(),
496 plan_column.source_name(),
497 )),
498 ),
499 None => diagnostics.push(
500 Diagnostic::error(
501 DiagnosticCode::SchemaMismatch,
502 "bulk target decimal precision/scale metadata is not available",
503 )
504 .with_field(FieldRef::new(
505 plan_column.source_index(),
506 plan_column.source_name(),
507 )),
508 ),
509 }
510 }
511}
512
513fn expected_direct_bulk_column_type(column: &DirectColumnPlan) -> Option<tiberius::ColumnType> {
514 match column.encoding() {
515 DirectColumnEncoding::Primitive(PrimitiveArrowToMssql::BooleanToBit) => {
516 if column.nullable() {
517 Some(tiberius::ColumnType::Bitn)
518 } else {
519 Some(tiberius::ColumnType::Bit)
520 }
521 }
522 DirectColumnEncoding::Primitive(PrimitiveArrowToMssql::UInt8ToTinyInt) => {
523 Some(tiberius::ColumnType::Int1)
524 }
525 DirectColumnEncoding::Primitive(
526 PrimitiveArrowToMssql::Int8ToSmallInt | PrimitiveArrowToMssql::Int16ToSmallInt,
527 ) => Some(tiberius::ColumnType::Int2),
528 DirectColumnEncoding::Primitive(PrimitiveArrowToMssql::Int32ToInt) => {
529 Some(tiberius::ColumnType::Int4)
530 }
531 DirectColumnEncoding::Primitive(PrimitiveArrowToMssql::UInt16ToInt) => {
532 Some(tiberius::ColumnType::Int4)
533 }
534 DirectColumnEncoding::Primitive(PrimitiveArrowToMssql::Int64ToBigInt) => {
535 Some(tiberius::ColumnType::Int8)
536 }
537 DirectColumnEncoding::Primitive(PrimitiveArrowToMssql::UInt32ToBigInt) => {
538 Some(tiberius::ColumnType::Int8)
539 }
540 DirectColumnEncoding::Primitive(PrimitiveArrowToMssql::UInt64ToCheckedBigInt) => {
541 Some(tiberius::ColumnType::Int8)
542 }
543 DirectColumnEncoding::Primitive(
544 PrimitiveArrowToMssql::Float16ToReal | PrimitiveArrowToMssql::Float32ToReal,
545 ) => Some(tiberius::ColumnType::Float4),
546 DirectColumnEncoding::Primitive(PrimitiveArrowToMssql::Float64ToFloat) => {
547 Some(tiberius::ColumnType::Float8)
548 }
549 DirectColumnEncoding::UInt64Decimal20_0 | DirectColumnEncoding::Decimal(_) => {
550 Some(tiberius::ColumnType::Decimaln)
551 }
552 DirectColumnEncoding::VariableWidth(VariableWidthArrowToMssql::StringToNVarChar {
553 ..
554 }) => Some(tiberius::ColumnType::NVarchar),
555 DirectColumnEncoding::VariableWidth(VariableWidthArrowToMssql::StringToAsciiVarChar {
556 ..
557 }) => Some(tiberius::ColumnType::BigVarChar),
558 DirectColumnEncoding::VariableWidth(VariableWidthArrowToMssql::BytesToVarBinary {
559 ..
560 }) => Some(tiberius::ColumnType::BigVarBin),
561 DirectColumnEncoding::FixedSizeBinary(
562 FixedSizeBinaryArrowToMssql::FixedSizeBinaryToBinary { .. },
563 ) => Some(tiberius::ColumnType::BigBinary),
564 DirectColumnEncoding::Temporal(TemporalArrowToMssql::Date32ToDate) => {
565 Some(tiberius::ColumnType::Daten)
566 }
567 DirectColumnEncoding::Temporal(TemporalArrowToMssql::Date64ToDateTime2) => {
568 Some(tiberius::ColumnType::Datetime2)
569 }
570 DirectColumnEncoding::Temporal(
571 TemporalArrowToMssql::TimestampSecondToDateTime2
572 | TemporalArrowToMssql::TimestampMillisecondToDateTime2
573 | TemporalArrowToMssql::TimestampMicrosecondToDateTime2
574 | TemporalArrowToMssql::TimestampNanosecondToDateTime2
575 | TemporalArrowToMssql::TimestampSecondTzToDateTime2
576 | TemporalArrowToMssql::TimestampMillisecondTzToDateTime2
577 | TemporalArrowToMssql::TimestampMicrosecondTzToDateTime2
578 | TemporalArrowToMssql::TimestampNanosecondTzToDateTime2,
579 ) => Some(tiberius::ColumnType::Datetime2),
580 DirectColumnEncoding::Temporal(
581 TemporalArrowToMssql::TimestampSecondToDateTime
582 | TemporalArrowToMssql::TimestampMillisecondToDateTime
583 | TemporalArrowToMssql::TimestampMicrosecondToDateTime
584 | TemporalArrowToMssql::TimestampNanosecondToDateTime
585 | TemporalArrowToMssql::TimestampSecondTzToDateTime
586 | TemporalArrowToMssql::TimestampMillisecondTzToDateTime
587 | TemporalArrowToMssql::TimestampMicrosecondTzToDateTime
588 | TemporalArrowToMssql::TimestampNanosecondTzToDateTime,
589 ) => Some(tiberius::ColumnType::Datetimen),
590 DirectColumnEncoding::Temporal(
591 TemporalArrowToMssql::Time32SecondToTime
592 | TemporalArrowToMssql::Time32MillisecondToTime
593 | TemporalArrowToMssql::Time64MicrosecondToTime
594 | TemporalArrowToMssql::Time64NanosecondToTime,
595 ) => Some(tiberius::ColumnType::Timen),
596 DirectColumnEncoding::Temporal(
597 TemporalArrowToMssql::TimestampSecondTzToDateTimeOffset
598 | TemporalArrowToMssql::TimestampMillisecondTzToDateTimeOffset
599 | TemporalArrowToMssql::TimestampMicrosecondTzToDateTimeOffset
600 | TemporalArrowToMssql::TimestampNanosecondTzToDateTimeOffset,
601 ) => Some(tiberius::ColumnType::DatetimeOffsetn),
602 }
603}
604
605fn expected_direct_decimal_precision_scale(column: &DirectColumnPlan) -> Option<(u8, u8)> {
606 match column.encoding() {
607 DirectColumnEncoding::UInt64Decimal20_0 => Some((20, 0)),
608 DirectColumnEncoding::Decimal(classification) => Some((
609 classification.target_precision(),
610 classification.target_scale(),
611 )),
612 _ => None,
613 }
614}
615
616fn bulk_target_column_diagnostic(
617 mapping: &SchemaMapping,
618 message: impl Into<String>,
619) -> Diagnostic {
620 Diagnostic::error(DiagnosticCode::SchemaMismatch, message).with_field(FieldRef::new(
621 mapping.arrow().index(),
622 mapping.arrow().name(),
623 ))
624}
625
626trait BulkTargetColumnMetadata {
627 fn ordinal(&self) -> usize;
628
629 fn name(&self) -> &str;
630
631 fn is_nullable(&self) -> bool;
632
633 fn column_type(&self) -> tiberius::ColumnType;
634
635 fn decimal_precision_scale(&self) -> Option<(u8, u8)> {
636 None
637 }
638}
639
640impl<T> BulkTargetColumnMetadata for &T
641where
642 T: BulkTargetColumnMetadata + ?Sized,
643{
644 fn ordinal(&self) -> usize {
645 (*self).ordinal()
646 }
647
648 fn name(&self) -> &str {
649 (*self).name()
650 }
651
652 fn is_nullable(&self) -> bool {
653 (*self).is_nullable()
654 }
655
656 fn column_type(&self) -> tiberius::ColumnType {
657 (*self).column_type()
658 }
659
660 fn decimal_precision_scale(&self) -> Option<(u8, u8)> {
661 (*self).decimal_precision_scale()
662 }
663}
664
665impl BulkTargetColumnMetadata for tiberius::BulkLoadColumn<'_> {
666 fn ordinal(&self) -> usize {
667 self.ordinal()
668 }
669
670 fn name(&self) -> &str {
671 self.name()
672 }
673
674 fn is_nullable(&self) -> bool {
675 self.is_nullable()
676 }
677
678 fn column_type(&self) -> tiberius::ColumnType {
679 self.column_type()
680 }
681
682 fn decimal_precision_scale(&self) -> Option<(u8, u8)> {
683 match self.type_info() {
684 tiberius::TypeInfo::VarLenSizedPrecision {
685 ty: tiberius::VarLenType::Decimaln | tiberius::VarLenType::Numericn,
686 precision,
687 scale,
688 ..
689 } => Some((*precision, *scale)),
690 _ => None,
691 }
692 }
693}
694
695async fn write_batch_to_sink<Sink>(
697 state: &mut WriterState,
698 sink: &mut Sink,
699 batch: &RecordBatch,
700) -> Result<WriteStats>
701where
702 Sink: TokenRowSink,
703{
704 let view = match record_batch_view(
705 batch,
706 state.mappings(),
707 state.schema_check(),
708 state.runtime_context(),
709 ) {
710 Ok(view) => view,
711 Err(err) => return Err(err.with_write_phase(WritePhase::BatchSchemaValidation)),
712 };
713 if let Err(err) = validate_batch_rows(&view) {
714 return Err(err.with_write_phase(WritePhase::ValueConversion));
715 }
716 let rows_written = usize_to_u64_saturating(view.row_count());
717
718 for row_index in 0..view.row_count() {
719 let row = match tiberius_row_owned(&view, row_index) {
720 Ok(row) => row,
721 Err(err) => return Err(err.with_write_phase(WritePhase::ValueConversion)),
722 };
723 if let Err(err) = sink.send_token_row(row).await {
724 return Err(err.with_write_phase(WritePhase::PacketWrite));
725 }
726 }
727
728 let stats = state.record_accepted_batch(rows_written);
729 Ok(stats)
730}
731
732async fn write_traced_batch_to_sink<Sink>(
734 state: &mut WriterState,
735 sink: &mut Sink,
736 batch: &RecordBatch,
737) -> Result<WriteStats>
738where
739 Sink: TokenRowSink,
740{
741 let trace = BatchWriteTrace::new(state.backend(), state.stats(), batch);
742 trace
743 .trace_result(write_batch_to_sink(state, sink, batch))
744 .await
745}
746
747trait TokenRowSink {
748 async fn send_token_row(&mut self, row: tiberius::TokenRow<'static>) -> Result<()>;
749}
750
751impl<S> TokenRowSink for tiberius::BulkLoadRequest<'_, S>
752where
753 S: AsyncRead + AsyncWrite + Unpin + Send,
754{
755 async fn send_token_row(&mut self, row: tiberius::TokenRow<'static>) -> Result<()> {
756 self.send(row)
757 .await
758 .map_err(|source| crate::Error::Tiberius { source })
759 }
760}
761
762#[cfg(test)]
764async fn write_direct_batch_to_sink<Sink>(
765 state: &mut WriterState,
766 sink: &mut Sink,
767 batch: &RecordBatch,
768) -> Result<WriteStats>
769where
770 Sink: RawRowsSink,
771{
772 write_direct_batch_to_sink_with_observer(state, sink, batch, DirectRawBatchObserver::disabled())
773 .await
774}
775
776async fn write_traced_direct_batch_to_sink<Sink>(
778 state: &mut WriterState,
779 sink: &mut Sink,
780 batch: &RecordBatch,
781) -> Result<WriteStats>
782where
783 Sink: RawRowsSink,
784{
785 let trace = BatchWriteTrace::new(state.backend(), state.stats(), batch);
786 let direct_observer = DirectRawBatchObserver::enabled(state.backend());
787 trace
788 .trace_result(write_direct_batch_to_sink_with_observer(
789 state,
790 sink,
791 batch,
792 direct_observer,
793 ))
794 .await
795}
796
797async fn write_direct_batch_to_sink_with_observer<Sink>(
802 state: &mut WriterState,
803 sink: &mut Sink,
804 batch: &RecordBatch,
805 direct_observer: DirectRawBatchObserver,
806) -> Result<WriteStats>
807where
808 Sink: RawRowsSink,
809{
810 let encoder = state
811 .direct_encoder()
812 .ok_or_else(|| crate::Error::BackendUnavailable {
813 backend: state.backend(),
814 reason: "direct bulk encoder is not available for this writer".to_owned(),
815 })
816 .map_err(|err| err.with_write_phase(WritePhase::DirectEncoding))?;
817
818 let measure_start = std::time::Instant::now();
819 let measured = encoder.measure_batch(batch);
820 let measured =
821 match profile::record_elapsed(measure_start, profile::record_measure_batch, measured) {
822 Ok(measured) => measured,
823 Err(err) => {
824 let phase = write_phase_for_batch_error(&err);
825 direct_observer.record_failed(phase.as_str(), batch, None, &err);
826 return Err(err.with_write_phase(phase));
827 }
828 };
829 direct_observer.record_measured(&measured, measure_start.elapsed());
830 let rows_written = usize_to_u64_saturating(measured.row_count());
831
832 let split_start = std::time::Instant::now();
833 let ranges = measured.row_ranges(DIRECT_RAW_MAX_PAYLOAD_BYTES);
834 let ranges = match profile::record_elapsed(split_start, profile::record_row_range_split, ranges)
835 {
836 Ok(ranges) => ranges,
837 Err(err) => {
838 direct_observer.record_failed(DIRECT_ENCODING_PHASE, batch, None, &err);
839 return Err(err.with_write_phase(WritePhase::DirectEncoding));
840 }
841 };
842 direct_observer.record_ranges_planned(&measured, &ranges, split_start.elapsed());
843
844 for range in ranges {
845 if let Err(err) = sink
846 .send_measured_raw_rows(encoder, batch, &measured, range, direct_observer)
847 .await
848 {
849 let phase = write_phase_for_batch_error(&err);
850 direct_observer.record_failed(phase.as_str(), batch, Some(range), &err);
851 return Err(err.with_write_phase(phase));
852 }
853 }
854
855 profile::record_accepted_batch(measured.row_count());
856 let stats = state.record_accepted_batch(rows_written);
857 Ok(stats)
858}
859
860trait RawRowsSink {
861 async fn send_measured_raw_rows(
862 &mut self,
863 encoder: &DirectEncoder,
864 batch: &RecordBatch,
865 measured: &MeasuredDirectBatch,
866 range: MeasuredRowRange,
867 direct_observer: DirectRawBatchObserver,
868 ) -> Result<()>;
869}
870
871impl<S> RawRowsSink for tiberius::BulkLoadRequest<'_, S>
872where
873 S: AsyncRead + AsyncWrite + Unpin + Send,
874{
875 async fn send_measured_raw_rows(
876 &mut self,
877 encoder: &DirectEncoder,
878 batch: &RecordBatch,
879 measured: &MeasuredDirectBatch,
880 range: MeasuredRowRange,
881 direct_observer: DirectRawBatchObserver,
882 ) -> Result<()> {
883 let encoded_bytes = measured.range_payload_len(range.start, range.len)?;
884 profile::record_row_range(encoded_bytes);
885
886 if !encoder.has_variable_width_column() {
887 let encode_start = std::time::Instant::now();
888 let payload =
889 encoder.encode_measured_batch_range(batch, measured, range.start, range.len)?;
890 profile::record_append_encode(encode_start.elapsed());
891
892 let send_start = std::time::Instant::now();
893 let send_result = self
894 .send_raw_rows_payload_checked(payload.bytes(), payload.row_token_offsets())
895 .await
896 .map_err(|source| crate::Error::Tiberius { source });
897 profile::record_send_total(send_start.elapsed());
898 if send_result.is_ok() {
899 direct_observer.record_packet_write_completed(
900 range,
901 payload.row_count(),
902 payload.bytes().len(),
903 send_start.elapsed(),
904 );
905 }
906 return send_result;
907 }
908
909 let mut encode_error = None;
910 let send_start = std::time::Instant::now();
911 let send_result = self
912 .send_raw_rows_with(|buf| {
913 let encode_start = std::time::Instant::now();
914 let encoded = encoder.encode_measured_batch_range_into(
915 batch,
916 measured,
917 range.start,
918 range.len,
919 buf,
920 );
921 profile::record_append_encode(encode_start.elapsed());
922
923 match encoded {
924 Ok(append) => Ok(append),
925 Err(err) => {
926 encode_error = Some(err);
927 Err(tiberius::error::Error::BulkInput(Cow::Borrowed(
928 "direct raw row encoding failed",
929 )))
930 }
931 }
932 })
933 .await;
934 profile::record_send_total(send_start.elapsed());
935
936 if let Some(err) = encode_error {
937 return Err(err);
938 }
939
940 let send_result = send_result.map_err(|source| crate::Error::Tiberius { source });
941 if send_result.is_ok() {
942 direct_observer.record_packet_write_completed(
943 range,
944 range.len,
945 encoded_bytes,
946 send_start.elapsed(),
947 );
948 }
949
950 send_result
951 }
952}
953
954fn resolve_backend(requested_backend: WriteBackend) -> Result<WriteBackend> {
955 match requested_backend {
956 WriteBackend::Auto | WriteBackend::DirectRawBulk => Ok(WriteBackend::DirectRawBulk),
957 WriteBackend::BaselineTokenRow => Ok(WriteBackend::BaselineTokenRow),
958 WriteBackend::DirectFramedBulk => Ok(WriteBackend::DirectFramedBulk),
959 }
960}
961
962fn execution_unavailable(backend: WriteBackend) -> crate::Error {
963 crate::Error::BackendUnavailable {
964 backend,
965 reason: "bulk writer execution is not implemented yet".to_owned(),
966 }
967}
968
969fn write_phase_for_batch_error(error: &crate::Error) -> WritePhase {
970 match error {
971 crate::Error::WritePhaseContext { phase, .. } => *phase,
972 crate::Error::ValueConversion { diagnostics }
973 if diagnostics
974 .all()
975 .iter()
976 .all(|diagnostic| diagnostic.code() == DiagnosticCode::SchemaMismatch) =>
977 {
978 WritePhase::BatchSchemaValidation
979 }
980 crate::Error::ValueConversion { .. } => WritePhase::ValueConversion,
981 crate::Error::DirectEncoding { .. } | crate::Error::BackendUnavailable { .. } => {
982 WritePhase::DirectEncoding
983 }
984 crate::Error::Tiberius { .. } => WritePhase::PacketWrite,
985 _ => WritePhase::BatchWrite,
986 }
987}
988
989fn usize_to_u64_saturating(value: usize) -> u64 {
990 u64::try_from(value).unwrap_or(u64::MAX)
991}
992
993#[cfg(test)]
994mod tests {
995 use std::{
996 borrow::Cow,
997 future::Future,
998 pin::Pin,
999 sync::{Arc, Mutex, MutexGuard},
1000 task::{Context, Poll, Waker},
1001 };
1002
1003 use arrow_array::{
1004 BinaryArray, Float64Array, Int32Array, RecordBatch, TimestampMicrosecondArray, UInt64Array,
1005 };
1006 use arrow_schema::{DataType, Field, Schema, TimeUnit};
1007 use futures_util::io::{AsyncRead, AsyncWrite};
1008
1009 use super::{
1010 BulkTargetColumnMetadata, DIRECT_RAW_MAX_PAYLOAD_BYTES, DirectEncoder, MeasuredDirectBatch,
1011 MeasuredRowRange, RawRowsSink, TokenRowSink, WriteBackend, WriteOptions, WriteStats,
1012 WriterState, bulk_insert_table_sql, record_batch_view, resolve_backend, tiberius_row_owned,
1013 validate_batch_rows, validate_bulk_target_columns,
1014 validate_direct_bulk_target_column_types, write_batch_to_sink, write_direct_batch_to_sink,
1015 };
1016 use crate::observability::writer::DirectRawBatchObserver;
1017 use crate::write::context::RuntimeConversionContext;
1018 use crate::{
1019 ArrowFieldRef, DiagnosticCode, Error, Identifier, MssqlColumn, MssqlProfile, MssqlType,
1020 MssqlTypeLength, NanosecondPolicy, PlanOptions, PlannedSchema, SchemaCheck, SchemaMapping,
1021 TableName, TimestampPolicy, WritePhase,
1022 };
1023
1024 static DIRECT_RAW_TRACE_TEST_LOCK: Mutex<()> = Mutex::new(());
1025
1026 fn direct_raw_trace_test_guard() -> MutexGuard<'static, ()> {
1027 match DIRECT_RAW_TRACE_TEST_LOCK.lock() {
1028 Ok(guard) => guard,
1029 Err(poisoned) => poisoned.into_inner(),
1030 }
1031 }
1032
1033 #[test]
1034 fn write_backend_defaults_to_auto() {
1035 assert_eq!(WriteBackend::default(), WriteBackend::Auto);
1036 }
1037
1038 #[test]
1039 fn write_options_default_to_auto_backend_and_strict_schema_check() {
1040 let options = WriteOptions::default();
1041
1042 assert_eq!(options.backend, WriteBackend::Auto);
1043 assert_eq!(options.schema_check, SchemaCheck::Strict);
1044 }
1045
1046 #[test]
1047 fn write_options_preserve_explicit_backend_selection() {
1048 for backend in [
1049 WriteBackend::Auto,
1050 WriteBackend::BaselineTokenRow,
1051 WriteBackend::DirectFramedBulk,
1052 WriteBackend::DirectRawBulk,
1053 ] {
1054 let options = WriteOptions {
1055 backend,
1056 schema_check: SchemaCheck::Strict,
1057 };
1058
1059 assert_eq!(options.backend, backend);
1060 assert_eq!(options.schema_check, SchemaCheck::Strict);
1061 }
1062 }
1063
1064 #[test]
1065 fn write_stats_default_to_zero() {
1066 let stats = WriteStats::default();
1067
1068 assert_eq!(stats.rows_written, 0);
1069 assert_eq!(stats.batches_written, 0);
1070 }
1071
1072 #[test]
1073 fn auto_backend_resolves_to_direct_raw_bulk() {
1074 assert_eq!(
1075 resolve_backend(WriteBackend::Auto).unwrap(),
1076 WriteBackend::DirectRawBulk
1077 );
1078 }
1079
1080 #[test]
1081 fn explicit_backends_resolve_to_requested_backend() {
1082 assert_eq!(
1083 resolve_backend(WriteBackend::BaselineTokenRow).unwrap(),
1084 WriteBackend::BaselineTokenRow
1085 );
1086 assert_eq!(
1087 resolve_backend(WriteBackend::DirectFramedBulk).unwrap(),
1088 WriteBackend::DirectFramedBulk
1089 );
1090 assert_eq!(
1091 resolve_backend(WriteBackend::DirectRawBulk).unwrap(),
1092 WriteBackend::DirectRawBulk
1093 );
1094 }
1095
1096 #[test]
1097 fn writer_state_starts_with_resolved_backend_mappings_and_zero_stats() {
1098 let mappings = vec![mapping("id")];
1099
1100 let state = WriterState::new(
1101 WriteBackend::Auto,
1102 SchemaCheck::Strict,
1103 planned_schema(mappings.clone()),
1104 )
1105 .unwrap();
1106
1107 assert_eq!(state.backend(), WriteBackend::DirectRawBulk);
1108 assert!(state.direct_encoder().is_some());
1109 assert_eq!(state.schema_check(), SchemaCheck::Strict);
1110 assert_eq!(state.mappings(), mappings.as_slice());
1111 assert_eq!(
1112 state.runtime_context().plan_options(),
1113 PlanOptions::default()
1114 );
1115 assert_eq!(state.stats(), WriteStats::default());
1116 }
1117
1118 #[test]
1119 fn writer_state_uses_runtime_context_from_planned_schema() {
1120 let profile = MssqlProfile::sql_server_2017_compat_140();
1121 let plan_options = PlanOptions {
1122 nanosecond_policy: NanosecondPolicy::TruncateTo100ns,
1123 ..PlanOptions::default()
1124 };
1125 let state = WriterState::new(
1126 WriteBackend::BaselineTokenRow,
1127 SchemaCheck::Strict,
1128 planned_schema_with_profile_and_options(profile, plan_options, vec![mapping("id")]),
1129 )
1130 .unwrap();
1131
1132 assert_eq!(state.runtime_context().profile(), profile);
1133 assert_eq!(state.runtime_context().plan_options(), plan_options);
1134 assert_eq!(
1135 state.runtime_context().nanosecond_policy(),
1136 NanosecondPolicy::TruncateTo100ns
1137 );
1138 }
1139
1140 #[test]
1141 fn direct_writer_state_builds_encoder_for_supported_mappings() {
1142 let mappings = vec![
1143 mapping("id32"),
1144 SchemaMapping::new(
1145 ArrowFieldRef::new(1, "id64".to_owned(), false, DataType::Int64),
1146 MssqlColumn::new(Identifier::new("id64").unwrap(), MssqlType::BigInt, false),
1147 ),
1148 float_mapping_at(2, "score"),
1149 SchemaMapping::new(
1150 ArrowFieldRef::new(3, "name".to_owned(), true, DataType::Utf8),
1151 MssqlColumn::new(
1152 Identifier::new("name").unwrap(),
1153 MssqlType::NVarChar(crate::MssqlTypeLength::Max),
1154 true,
1155 ),
1156 ),
1157 ];
1158
1159 for backend in [WriteBackend::DirectFramedBulk, WriteBackend::DirectRawBulk] {
1160 let state = WriterState::new(
1161 backend,
1162 SchemaCheck::Strict,
1163 planned_schema(mappings.clone()),
1164 )
1165 .unwrap();
1166
1167 assert_eq!(state.backend(), backend);
1168 assert!(state.direct_encoder().is_some());
1169 }
1170 }
1171
1172 #[test]
1173 fn direct_writer_state_rejects_unsupported_mappings() {
1174 let mappings = vec![SchemaMapping::new(
1175 ArrowFieldRef::new(
1176 0,
1177 "list_value".to_owned(),
1178 true,
1179 DataType::List(Arc::new(Field::new("item", DataType::Int32, true))),
1180 ),
1181 MssqlColumn::new(
1182 Identifier::new("list_value").unwrap(),
1183 MssqlType::NVarChar(MssqlTypeLength::Max),
1184 true,
1185 ),
1186 )];
1187
1188 let err = WriterState::new(
1189 WriteBackend::DirectRawBulk,
1190 SchemaCheck::Strict,
1191 planned_schema(mappings),
1192 )
1193 .unwrap_err();
1194
1195 let Error::DirectEncoding { diagnostics } = err else {
1196 panic!("expected direct encoding error");
1197 };
1198 assert_eq!(diagnostics.len(), 1);
1199 assert_eq!(
1200 diagnostics.all()[0].code(),
1201 DiagnosticCode::DirectEncodingUnsupportedMapping
1202 );
1203 }
1204
1205 #[test]
1206 fn writer_state_accumulates_accepted_batch_stats() {
1207 let mut state = WriterState::new(
1208 WriteBackend::BaselineTokenRow,
1209 SchemaCheck::Strict,
1210 planned_schema(Vec::new()),
1211 )
1212 .unwrap();
1213
1214 assert_eq!(
1215 state.record_accepted_batch(0),
1216 WriteStats {
1217 rows_written: 0,
1218 batches_written: 1
1219 }
1220 );
1221 assert_eq!(
1222 state.record_accepted_batch(3),
1223 WriteStats {
1224 rows_written: 3,
1225 batches_written: 2
1226 }
1227 );
1228 assert_eq!(
1229 state.record_accepted_batch(5),
1230 WriteStats {
1231 rows_written: 8,
1232 batches_written: 3
1233 }
1234 );
1235 }
1236
1237 #[test]
1238 fn bulk_insert_table_sql_uses_quoted_table_name() {
1239 let table = TableName::new("dbo]x", "target.table").unwrap();
1240
1241 assert_eq!(bulk_insert_table_sql(&table), "[dbo]]x].[target.table]");
1242 }
1243
1244 #[test]
1245 fn strict_batch_validation_accepts_supported_rows_without_owning_payloads() {
1246 let batch = int32_batch("id", &[1, 2]);
1247 let mappings = [mapping("id")];
1248 let view = record_batch_view(
1249 &batch,
1250 &mappings,
1251 SchemaCheck::Strict,
1252 runtime_context_with_options(PlanOptions::default()),
1253 )
1254 .unwrap();
1255
1256 validate_batch_rows(&view).unwrap();
1257
1258 let row = tiberius_row_owned(&view, 1).unwrap();
1259 assert_eq!(row.get(0), Some(&tiberius::ColumnData::I32(Some(2))));
1260 }
1261
1262 #[test]
1263 fn strict_batch_view_rejects_runtime_schema_mismatch_before_send() {
1264 let batch = int32_batch("renamed_id", &[1]);
1265 let err = record_batch_view(
1266 &batch,
1267 &[mapping("id")],
1268 SchemaCheck::Strict,
1269 runtime_context_with_options(PlanOptions::default()),
1270 )
1271 .unwrap_err();
1272
1273 let Error::ValueConversion { diagnostics } = err else {
1274 panic!("expected value conversion error");
1275 };
1276 assert_eq!(diagnostics.len(), 1);
1277 let diagnostic = &diagnostics.all()[0];
1278 assert_eq!(diagnostic.code(), DiagnosticCode::SchemaMismatch);
1279 assert_eq!(diagnostic.field().map(|field| field.name()), Some("id"));
1280 }
1281
1282 #[test]
1283 fn strict_batch_validation_rejects_bad_later_row_before_any_send() {
1284 let schema = Arc::new(Schema::new(vec![Field::new(
1285 "amount",
1286 DataType::Float64,
1287 false,
1288 )]));
1289 let batch = RecordBatch::try_new(
1290 schema,
1291 vec![Arc::new(Float64Array::from(vec![
1292 Some(1.0),
1293 Some(f64::NAN),
1294 ]))],
1295 )
1296 .unwrap();
1297 let mappings = [SchemaMapping::new(
1298 ArrowFieldRef::new(0, "amount".to_owned(), false, DataType::Float64),
1299 MssqlColumn::new(
1300 Identifier::new("amount").unwrap(),
1301 MssqlType::Float { precision: 53 },
1302 false,
1303 ),
1304 )];
1305
1306 let view = record_batch_view(
1307 &batch,
1308 &mappings,
1309 SchemaCheck::Strict,
1310 runtime_context_with_options(PlanOptions::default()),
1311 )
1312 .unwrap();
1313 let err = validate_batch_rows(&view).unwrap_err();
1314
1315 let Error::ValueConversion { diagnostics } = err else {
1316 panic!("expected value conversion error");
1317 };
1318 assert_eq!(diagnostics.len(), 1);
1319 let diagnostic = &diagnostics.all()[0];
1320 assert_eq!(diagnostic.code(), DiagnosticCode::NonFiniteFloat);
1321 assert_eq!(diagnostic.row(), Some(1));
1322 }
1323
1324 #[test]
1325 fn bulk_target_column_validation_accepts_matching_metadata() {
1326 let mappings = vec![mapping("id")];
1327 let columns = vec![bulk_target_column(0, "id", false)];
1328
1329 validate_bulk_target_columns(columns.into_iter(), &mappings).unwrap();
1330 }
1331
1332 #[test]
1333 fn bulk_target_column_validation_rejects_missing_target_columns() {
1334 let mappings = vec![mapping("id")];
1335 let columns = Vec::<FakeBulkTargetColumn>::new();
1336
1337 let err = validate_bulk_target_columns(columns.into_iter(), &mappings).unwrap_err();
1338
1339 let Error::ValueConversion { diagnostics } = err else {
1340 panic!("expected value conversion error");
1341 };
1342 assert_eq!(diagnostics.len(), 1);
1343 assert_eq!(diagnostics.all()[0].code(), DiagnosticCode::SchemaMismatch);
1344 assert_eq!(
1345 diagnostics.all()[0].message(),
1346 "bulk target has 0 updateable column(s) but mappings contain 1 column(s)"
1347 );
1348 }
1349
1350 #[test]
1351 fn bulk_target_column_validation_rejects_ordinal_name_and_nullability_drift() {
1352 let mappings = vec![mapping("id")];
1353 let columns = vec![bulk_target_column(7, "id]; DROP TABLE target;--", true)];
1354
1355 let err = validate_bulk_target_columns(columns.into_iter(), &mappings).unwrap_err();
1356
1357 let Error::ValueConversion { diagnostics } = err else {
1358 panic!("expected value conversion error");
1359 };
1360 assert_eq!(diagnostics.len(), 3);
1361 assert!(
1362 diagnostics
1363 .all()
1364 .iter()
1365 .all(|diagnostic| diagnostic.code() == DiagnosticCode::SchemaMismatch)
1366 );
1367 assert!(
1368 diagnostics
1369 .all()
1370 .iter()
1371 .all(|diagnostic| diagnostic.field().map(|field| field.name()) == Some("id"))
1372 );
1373 assert!(
1374 diagnostics
1375 .all()
1376 .iter()
1377 .any(|diagnostic| diagnostic.message().contains("ordinal 7"))
1378 );
1379 assert!(
1380 diagnostics
1381 .all()
1382 .iter()
1383 .any(|diagnostic| diagnostic.message().contains("DROP TABLE"))
1384 );
1385 assert!(
1386 diagnostics
1387 .all()
1388 .iter()
1389 .any(|diagnostic| diagnostic.message().contains("nullability true"))
1390 );
1391 }
1392
1393 #[test]
1394 fn direct_bulk_target_type_validation_accepts_matching_primitive_metadata() {
1395 let mappings = vec![mapping("id")];
1396 let state = WriterState::new(
1397 WriteBackend::DirectRawBulk,
1398 SchemaCheck::Strict,
1399 planned_schema(mappings),
1400 )
1401 .unwrap();
1402 let columns = vec![bulk_target_column_with_type(
1403 0,
1404 "id",
1405 false,
1406 tiberius::ColumnType::Int4,
1407 )];
1408
1409 validate_direct_bulk_target_column_types(
1410 columns.into_iter(),
1411 state.direct_encoder().unwrap().plan(),
1412 )
1413 .unwrap();
1414 }
1415
1416 #[test]
1417 fn direct_bulk_target_type_validation_accepts_issue_75_integer_metadata() {
1418 let mappings = vec![
1419 schema_mapping_at(0, "tiny", DataType::UInt8, MssqlType::TinyInt, false),
1420 schema_mapping_at(1, "signed_tiny", DataType::Int8, MssqlType::SmallInt, false),
1421 schema_mapping_at(2, "small", DataType::Int16, MssqlType::SmallInt, false),
1422 schema_mapping_at(
1423 3,
1424 "unsigned_medium",
1425 DataType::UInt16,
1426 MssqlType::Int,
1427 false,
1428 ),
1429 schema_mapping_at(
1430 4,
1431 "unsigned_total",
1432 DataType::UInt32,
1433 MssqlType::BigInt,
1434 false,
1435 ),
1436 ];
1437 let state = WriterState::new(
1438 WriteBackend::DirectRawBulk,
1439 SchemaCheck::Strict,
1440 planned_schema(mappings),
1441 )
1442 .unwrap();
1443 let columns = vec![
1444 bulk_target_column_with_type(0, "tiny", false, tiberius::ColumnType::Int1),
1445 bulk_target_column_with_type(1, "signed_tiny", false, tiberius::ColumnType::Int2),
1446 bulk_target_column_with_type(2, "small", false, tiberius::ColumnType::Int2),
1447 bulk_target_column_with_type(3, "unsigned_medium", false, tiberius::ColumnType::Int4),
1448 bulk_target_column_with_type(4, "unsigned_total", false, tiberius::ColumnType::Int8),
1449 ];
1450
1451 validate_direct_bulk_target_column_types(
1452 columns.into_iter(),
1453 state.direct_encoder().unwrap().plan(),
1454 )
1455 .unwrap();
1456 }
1457
1458 #[test]
1459 fn direct_bulk_target_type_validation_accepts_issue_75_float32_metadata() {
1460 let mappings = vec![schema_mapping_at(
1461 0,
1462 "real_value",
1463 DataType::Float32,
1464 MssqlType::Real,
1465 false,
1466 )];
1467 let state = WriterState::new(
1468 WriteBackend::DirectRawBulk,
1469 SchemaCheck::Strict,
1470 planned_schema(mappings),
1471 )
1472 .unwrap();
1473 let columns = vec![bulk_target_column_with_type(
1474 0,
1475 "real_value",
1476 false,
1477 tiberius::ColumnType::Float4,
1478 )];
1479
1480 validate_direct_bulk_target_column_types(
1481 columns.into_iter(),
1482 state.direct_encoder().unwrap().plan(),
1483 )
1484 .unwrap();
1485 }
1486
1487 #[test]
1488 fn direct_bulk_target_type_validation_accepts_uint64_policy_metadata() {
1489 let mappings = vec![
1490 schema_mapping_at(0, "checked", DataType::UInt64, MssqlType::BigInt, false),
1491 schema_mapping_at(
1492 1,
1493 "decimal",
1494 DataType::UInt64,
1495 MssqlType::Decimal {
1496 precision: 20,
1497 scale: 0,
1498 },
1499 false,
1500 ),
1501 ];
1502 let state = WriterState::new(
1503 WriteBackend::DirectRawBulk,
1504 SchemaCheck::Strict,
1505 planned_schema(mappings),
1506 )
1507 .unwrap();
1508 let columns = vec![
1509 bulk_target_column_with_type(0, "checked", false, tiberius::ColumnType::Int8),
1510 bulk_target_decimal_column(1, "decimal", false, 20, 0),
1511 ];
1512
1513 validate_direct_bulk_target_column_types(
1514 columns.into_iter(),
1515 state.direct_encoder().unwrap().plan(),
1516 )
1517 .unwrap();
1518 }
1519
1520 #[test]
1521 fn direct_bulk_target_type_validation_rejects_uint64_decimal_precision_drift() {
1522 let mappings = vec![schema_mapping_at(
1523 0,
1524 "decimal",
1525 DataType::UInt64,
1526 MssqlType::Decimal {
1527 precision: 20,
1528 scale: 0,
1529 },
1530 false,
1531 )];
1532 let state = WriterState::new(
1533 WriteBackend::DirectRawBulk,
1534 SchemaCheck::Strict,
1535 planned_schema(mappings),
1536 )
1537 .unwrap();
1538 let columns = vec![bulk_target_decimal_column(0, "decimal", false, 19, 0)];
1539
1540 let err = validate_direct_bulk_target_column_types(
1541 columns.into_iter(),
1542 state.direct_encoder().unwrap().plan(),
1543 )
1544 .unwrap_err();
1545
1546 let Error::ValueConversion { diagnostics } = err else {
1547 panic!("expected value conversion error");
1548 };
1549 assert_eq!(diagnostics.len(), 1);
1550 let diagnostic = &diagnostics.all()[0];
1551 assert_eq!(diagnostic.code(), DiagnosticCode::SchemaMismatch);
1552 assert!(diagnostic.message().contains("precision/scale (19,0)"));
1553 assert_eq!(
1554 diagnostic
1555 .field()
1556 .map(|field| (field.index(), field.name())),
1557 Some((0, "decimal"))
1558 );
1559 }
1560
1561 #[test]
1562 fn direct_bulk_target_type_validation_accepts_matching_variable_width_metadata() {
1563 let mappings = vec![
1564 utf8_mapping_at(0, "name"),
1565 schema_mapping_at(
1566 1,
1567 "ascii_code",
1568 DataType::Utf8,
1569 MssqlType::VarChar(MssqlTypeLength::Bounded(16)),
1570 false,
1571 ),
1572 binary_mapping_at(2, "payload"),
1573 ];
1574 let state = WriterState::new(
1575 WriteBackend::DirectRawBulk,
1576 SchemaCheck::Strict,
1577 planned_schema(mappings),
1578 )
1579 .unwrap();
1580 let columns = vec![
1581 bulk_target_column_with_type(0, "name", false, tiberius::ColumnType::NVarchar),
1582 bulk_target_column_with_type(1, "ascii_code", false, tiberius::ColumnType::BigVarChar),
1583 bulk_target_column_with_type(2, "payload", false, tiberius::ColumnType::BigVarBin),
1584 ];
1585
1586 validate_direct_bulk_target_column_types(
1587 columns.into_iter(),
1588 state.direct_encoder().unwrap().plan(),
1589 )
1590 .unwrap();
1591 }
1592
1593 #[test]
1594 fn direct_bulk_target_type_validation_accepts_matching_large_variable_width_metadata() {
1595 let mappings = vec![
1596 schema_mapping_at(
1597 0,
1598 "large_name",
1599 DataType::LargeUtf8,
1600 MssqlType::NVarChar(MssqlTypeLength::Max),
1601 false,
1602 ),
1603 schema_mapping_at(
1604 1,
1605 "large_payload",
1606 DataType::LargeBinary,
1607 MssqlType::VarBinary(MssqlTypeLength::Max),
1608 false,
1609 ),
1610 ];
1611 let state = WriterState::new(
1612 WriteBackend::DirectRawBulk,
1613 SchemaCheck::Strict,
1614 planned_schema(mappings),
1615 )
1616 .unwrap();
1617 let columns = vec![
1618 bulk_target_column_with_type(0, "large_name", false, tiberius::ColumnType::NVarchar),
1619 bulk_target_column_with_type(
1620 1,
1621 "large_payload",
1622 false,
1623 tiberius::ColumnType::BigVarBin,
1624 ),
1625 ];
1626
1627 validate_direct_bulk_target_column_types(
1628 columns.into_iter(),
1629 state.direct_encoder().unwrap().plan(),
1630 )
1631 .unwrap();
1632 }
1633
1634 #[test]
1635 fn direct_bulk_target_type_validation_accepts_fixed_size_binary_metadata() {
1636 let mappings = vec![fixed_size_binary_mapping_at(0, "digest", 32)];
1637 let state = WriterState::new(
1638 WriteBackend::DirectRawBulk,
1639 SchemaCheck::Strict,
1640 planned_schema(mappings),
1641 )
1642 .unwrap();
1643 let columns = vec![bulk_target_column_with_type(
1644 0,
1645 "digest",
1646 false,
1647 tiberius::ColumnType::BigBinary,
1648 )];
1649
1650 validate_direct_bulk_target_column_types(
1651 columns.into_iter(),
1652 state.direct_encoder().unwrap().plan(),
1653 )
1654 .unwrap();
1655 }
1656
1657 #[test]
1658 fn direct_bulk_target_type_validation_rejects_fixed_size_binary_as_varbinary() {
1659 let mappings = vec![fixed_size_binary_mapping_at(0, "digest", 32)];
1660 let state = WriterState::new(
1661 WriteBackend::DirectRawBulk,
1662 SchemaCheck::Strict,
1663 planned_schema(mappings),
1664 )
1665 .unwrap();
1666 let columns = vec![bulk_target_column_with_type(
1667 0,
1668 "digest",
1669 false,
1670 tiberius::ColumnType::BigVarBin,
1671 )];
1672
1673 let err = validate_direct_bulk_target_column_types(
1674 columns.into_iter(),
1675 state.direct_encoder().unwrap().plan(),
1676 )
1677 .unwrap_err();
1678
1679 let Error::ValueConversion { diagnostics } = err else {
1680 panic!("expected value conversion error");
1681 };
1682 assert_eq!(diagnostics.len(), 1);
1683 let diagnostic = &diagnostics.all()[0];
1684 assert_eq!(diagnostic.code(), DiagnosticCode::SchemaMismatch);
1685 assert_eq!(diagnostic.field().map(|field| field.name()), Some("digest"));
1686 assert!(diagnostic.message().contains(
1687 "bulk target column type BigVarBin does not match direct encoder type BigBinary"
1688 ));
1689 }
1690
1691 #[test]
1692 fn direct_bulk_target_type_validation_accepts_date_metadata() {
1693 let mappings = vec![
1694 SchemaMapping::new(
1695 ArrowFieldRef::new(0, "created_on".to_owned(), true, DataType::Date32),
1696 MssqlColumn::new(
1697 Identifier::new("created_on").unwrap(),
1698 MssqlType::Date,
1699 true,
1700 ),
1701 ),
1702 SchemaMapping::new(
1703 ArrowFieldRef::new(1, "created_at".to_owned(), true, DataType::Date64),
1704 MssqlColumn::new(
1705 Identifier::new("created_at").unwrap(),
1706 MssqlType::DateTime2 { precision: 3 },
1707 true,
1708 ),
1709 ),
1710 ];
1711 let state = WriterState::new(
1712 WriteBackend::DirectRawBulk,
1713 SchemaCheck::Strict,
1714 planned_schema(mappings),
1715 )
1716 .unwrap();
1717 let columns = vec![
1718 bulk_target_column_with_type(0, "created_on", true, tiberius::ColumnType::Daten),
1719 bulk_target_column_with_type(1, "created_at", true, tiberius::ColumnType::Datetime2),
1720 ];
1721
1722 validate_direct_bulk_target_column_types(
1723 columns.into_iter(),
1724 state.direct_encoder().unwrap().plan(),
1725 )
1726 .unwrap();
1727 }
1728
1729 #[test]
1730 fn direct_bulk_target_type_validation_accepts_datetime_metadata() {
1731 let mappings = vec![SchemaMapping::new(
1732 ArrowFieldRef::new(
1733 0,
1734 "created_at".to_owned(),
1735 false,
1736 DataType::Timestamp(TimeUnit::Microsecond, None),
1737 ),
1738 MssqlColumn::new(
1739 Identifier::new("created_at").unwrap(),
1740 MssqlType::DateTime,
1741 false,
1742 ),
1743 )];
1744 let state = WriterState::new(
1745 WriteBackend::DirectRawBulk,
1746 SchemaCheck::Strict,
1747 planned_schema(mappings),
1748 )
1749 .unwrap();
1750 for column_type in [
1751 tiberius::ColumnType::Datetime,
1752 tiberius::ColumnType::Datetimen,
1753 ] {
1754 let columns = vec![bulk_target_column_with_type(
1755 0,
1756 "created_at",
1757 false,
1758 column_type,
1759 )];
1760
1761 validate_direct_bulk_target_column_types(
1762 columns.into_iter(),
1763 state.direct_encoder().unwrap().plan(),
1764 )
1765 .unwrap();
1766 }
1767 }
1768
1769 #[test]
1770 fn direct_bulk_target_type_validation_rejects_variable_width_type_swap() {
1771 let mappings = vec![utf8_mapping_at(0, "name"), binary_mapping_at(1, "payload")];
1772 let state = WriterState::new(
1773 WriteBackend::DirectRawBulk,
1774 SchemaCheck::Strict,
1775 planned_schema(mappings),
1776 )
1777 .unwrap();
1778 let columns = vec![
1779 bulk_target_column_with_type(0, "name", false, tiberius::ColumnType::BigVarBin),
1780 bulk_target_column_with_type(1, "payload", false, tiberius::ColumnType::NVarchar),
1781 ];
1782
1783 let err = validate_direct_bulk_target_column_types(
1784 columns.into_iter(),
1785 state.direct_encoder().unwrap().plan(),
1786 )
1787 .unwrap_err();
1788
1789 let Error::ValueConversion { diagnostics } = err else {
1790 panic!("expected value conversion error");
1791 };
1792 assert_eq!(diagnostics.len(), 2);
1793 assert!(
1794 diagnostics
1795 .all()
1796 .iter()
1797 .any(|diagnostic| diagnostic.message().contains("NVarchar"))
1798 );
1799 assert!(
1800 diagnostics
1801 .all()
1802 .iter()
1803 .any(|diagnostic| diagnostic.message().contains("BigVarBin"))
1804 );
1805 }
1806
1807 #[test]
1808 fn direct_bulk_target_type_validation_rejects_same_name_with_wrong_type() {
1809 let mappings = vec![mapping("id")];
1810 let state = WriterState::new(
1811 WriteBackend::DirectRawBulk,
1812 SchemaCheck::Strict,
1813 planned_schema(mappings),
1814 )
1815 .unwrap();
1816 let columns = vec![bulk_target_column_with_type(
1817 0,
1818 "id",
1819 false,
1820 tiberius::ColumnType::Int8,
1821 )];
1822
1823 let err = validate_direct_bulk_target_column_types(
1824 columns.into_iter(),
1825 state.direct_encoder().unwrap().plan(),
1826 )
1827 .unwrap_err();
1828
1829 let Error::ValueConversion { diagnostics } = err else {
1830 panic!("expected value conversion error");
1831 };
1832 assert_eq!(diagnostics.len(), 1);
1833 let diagnostic = &diagnostics.all()[0];
1834 assert_eq!(diagnostic.code(), DiagnosticCode::SchemaMismatch);
1835 assert_eq!(diagnostic.field().map(|field| field.name()), Some("id"));
1836 assert!(
1837 diagnostic
1838 .message()
1839 .contains("bulk target column type Int8 does not match direct encoder type Int4")
1840 );
1841 }
1842
1843 #[test]
1844 fn write_batch_to_sink_accepts_empty_matching_batch() {
1845 let mappings = vec![mapping("id")];
1846 let mut state = WriterState::new(
1847 WriteBackend::BaselineTokenRow,
1848 SchemaCheck::Strict,
1849 planned_schema(mappings),
1850 )
1851 .unwrap();
1852 let mut sink = RecordingSink::default();
1853 let batch = int32_batch("id", &[]);
1854
1855 let stats = poll_ready(write_batch_to_sink(&mut state, &mut sink, &batch)).unwrap();
1856
1857 assert_eq!(
1858 stats,
1859 WriteStats {
1860 rows_written: 0,
1861 batches_written: 1
1862 }
1863 );
1864 assert!(sink.rows.is_empty());
1865 }
1866
1867 #[test]
1868 fn write_batch_to_sink_accumulates_multi_batch_stats() {
1869 let mappings = vec![mapping("id")];
1870 let mut state = WriterState::new(
1871 WriteBackend::BaselineTokenRow,
1872 SchemaCheck::Strict,
1873 planned_schema(mappings),
1874 )
1875 .unwrap();
1876 let mut sink = RecordingSink::default();
1877
1878 let first = poll_ready(write_batch_to_sink(
1879 &mut state,
1880 &mut sink,
1881 &int32_batch("id", &[10, 20]),
1882 ))
1883 .unwrap();
1884 let second = poll_ready(write_batch_to_sink(
1885 &mut state,
1886 &mut sink,
1887 &int32_batch("id", &[30]),
1888 ))
1889 .unwrap();
1890
1891 assert_eq!(
1892 first,
1893 WriteStats {
1894 rows_written: 2,
1895 batches_written: 1
1896 }
1897 );
1898 assert_eq!(
1899 second,
1900 WriteStats {
1901 rows_written: 3,
1902 batches_written: 2
1903 }
1904 );
1905 assert_eq!(sink.rows.len(), 3);
1906 assert_eq!(
1907 sink.rows[2].get(0),
1908 Some(&tiberius::ColumnData::I32(Some(30)))
1909 );
1910 }
1911
1912 #[test]
1913 fn write_batch_to_sink_sends_timestamp_datetime_cells() {
1914 let mappings = vec![SchemaMapping::new(
1915 ArrowFieldRef::new(
1916 0,
1917 "created_at".to_owned(),
1918 true,
1919 DataType::Timestamp(TimeUnit::Microsecond, None),
1920 ),
1921 MssqlColumn::new(
1922 Identifier::new("created_at").unwrap(),
1923 MssqlType::DateTime,
1924 true,
1925 ),
1926 )];
1927 let options = PlanOptions {
1928 timestamp_policy: TimestampPolicy::DateTime,
1929 ..PlanOptions::default()
1930 };
1931 let mut state = WriterState::new(
1932 WriteBackend::BaselineTokenRow,
1933 SchemaCheck::Strict,
1934 planned_schema_with_options(mappings, options),
1935 )
1936 .unwrap();
1937 let mut sink = RecordingSink::default();
1938 let batch =
1939 timestamp_microsecond_batch("created_at", &[Some(1_700), Some(86_399_999_000), None]);
1940
1941 let stats = poll_ready(write_batch_to_sink(&mut state, &mut sink, &batch)).unwrap();
1942
1943 assert_eq!(
1944 stats,
1945 WriteStats {
1946 rows_written: 3,
1947 batches_written: 1
1948 }
1949 );
1950 assert_eq!(sink.rows.len(), 3);
1951 assert_eq!(
1952 sink.rows[0].get(0),
1953 Some(&tiberius::ColumnData::DateTime(Some(
1954 tiberius::time::DateTime::new(25_567, 1)
1955 )))
1956 );
1957 assert_eq!(
1958 sink.rows[1].get(0),
1959 Some(&tiberius::ColumnData::DateTime(Some(
1960 tiberius::time::DateTime::new(25_568, 0)
1961 )))
1962 );
1963 assert_eq!(
1964 sink.rows[2].get(0),
1965 Some(&tiberius::ColumnData::DateTime(None))
1966 );
1967 }
1968
1969 #[test]
1970 fn write_batch_to_sink_conversion_failure_sends_nothing_and_keeps_stats() {
1971 let mappings = vec![float_mapping("amount")];
1972 let mut state = WriterState::new(
1973 WriteBackend::BaselineTokenRow,
1974 SchemaCheck::Strict,
1975 planned_schema(mappings),
1976 )
1977 .unwrap();
1978 let mut sink = RecordingSink::default();
1979 let batch = float64_batch("amount", &[Some(1.0), Some(f64::NAN)]);
1980
1981 let err = poll_ready(write_batch_to_sink(&mut state, &mut sink, &batch)).unwrap_err();
1982
1983 assert_write_phase(&err, WritePhase::ValueConversion);
1984 let Error::ValueConversion { diagnostics } = inner_error(&err) else {
1985 panic!("expected value conversion error");
1986 };
1987 assert_eq!(diagnostics.all()[0].code(), DiagnosticCode::NonFiniteFloat);
1988 assert_eq!(diagnostics.all()[0].row(), Some(1));
1989 assert!(sink.rows.is_empty());
1990 assert_eq!(state.stats(), WriteStats::default());
1991 }
1992
1993 #[test]
1994 fn write_batch_to_sink_send_failure_preserves_error_and_keeps_stats() {
1995 let mappings = vec![mapping("id")];
1996 let mut state = WriterState::new(
1997 WriteBackend::BaselineTokenRow,
1998 SchemaCheck::Strict,
1999 planned_schema(mappings),
2000 )
2001 .unwrap();
2002 let mut sink = RecordingSink {
2003 fail_on_send: Some(1),
2004 rows: Vec::new(),
2005 };
2006 let batch = int32_batch("id", &[1, 2, 3]);
2007
2008 let err = poll_ready(write_batch_to_sink(&mut state, &mut sink, &batch)).unwrap_err();
2009
2010 assert_write_phase(&err, WritePhase::PacketWrite);
2011 let Error::Tiberius { source } = inner_error(&err) else {
2012 panic!("expected tiberius error");
2013 };
2014 assert_eq!(
2015 source.to_string(),
2016 "BULK UPLOAD input failure: fake send failure"
2017 );
2018 assert_eq!(sink.rows.len(), 1);
2019 assert_eq!(state.stats(), WriteStats::default());
2020 }
2021
2022 #[test]
2023 fn write_direct_batch_to_sink_sends_one_checked_payload_per_batch() {
2024 let _trace_guard = direct_raw_trace_test_guard();
2025 let mappings = vec![mapping("id")];
2026 let mut state = WriterState::new(
2027 WriteBackend::DirectRawBulk,
2028 SchemaCheck::Strict,
2029 planned_schema(mappings),
2030 )
2031 .unwrap();
2032 let mut sink = RecordingRawSink::default();
2033 let batch = int32_batch("id", &[10, 20]);
2034
2035 let stats = poll_ready(write_direct_batch_to_sink(&mut state, &mut sink, &batch)).unwrap();
2036
2037 assert_eq!(
2038 stats,
2039 WriteStats {
2040 rows_written: 2,
2041 batches_written: 1
2042 }
2043 );
2044 assert_eq!(sink.payloads.len(), 1);
2045 assert_eq!(sink.payloads[0].row_token_offsets, vec![0, 5]);
2046 assert_eq!(
2047 sink.payloads[0].bytes,
2048 vec![0xD1, 10, 0, 0, 0, 0xD1, 20, 0, 0, 0]
2049 );
2050 }
2051
2052 #[test]
2053 fn write_direct_batch_to_sink_accumulates_multi_batch_stats() {
2054 let _trace_guard = direct_raw_trace_test_guard();
2055 let mappings = vec![mapping("id")];
2056 let mut state = WriterState::new(
2057 WriteBackend::DirectRawBulk,
2058 SchemaCheck::Strict,
2059 planned_schema(mappings),
2060 )
2061 .unwrap();
2062 let mut sink = RecordingRawSink::default();
2063
2064 let first = poll_ready(write_direct_batch_to_sink(
2065 &mut state,
2066 &mut sink,
2067 &int32_batch("id", &[10, 20]),
2068 ))
2069 .unwrap();
2070 let second = poll_ready(write_direct_batch_to_sink(
2071 &mut state,
2072 &mut sink,
2073 &int32_batch("id", &[30]),
2074 ))
2075 .unwrap();
2076
2077 assert_eq!(
2078 first,
2079 WriteStats {
2080 rows_written: 2,
2081 batches_written: 1
2082 }
2083 );
2084 assert_eq!(
2085 second,
2086 WriteStats {
2087 rows_written: 3,
2088 batches_written: 2
2089 }
2090 );
2091 assert_eq!(sink.payloads.len(), 2);
2092 assert_eq!(sink.payloads[1].bytes, vec![0xD1, 30, 0, 0, 0]);
2093 }
2094
2095 #[test]
2096 fn write_direct_batch_to_sink_chunks_measured_payloads_by_byte_limit() {
2097 let _trace_guard = direct_raw_trace_test_guard();
2098 let mappings = vec![binary_mapping_at(0, "payload")];
2099 let mut state = WriterState::new(
2100 WriteBackend::DirectRawBulk,
2101 SchemaCheck::Strict,
2102 planned_schema(mappings),
2103 )
2104 .unwrap();
2105 let mut sink = RecordingRawSink::default();
2106 let row_bytes = vec![0x5a; DIRECT_RAW_MAX_PAYLOAD_BYTES / 2 + 1];
2107 let batch = binary_batch("payload", &[row_bytes.as_slice(), row_bytes.as_slice()]);
2108
2109 let stats = poll_ready(write_direct_batch_to_sink(&mut state, &mut sink, &batch)).unwrap();
2110
2111 assert_eq!(
2112 stats,
2113 WriteStats {
2114 rows_written: 2,
2115 batches_written: 1
2116 }
2117 );
2118 assert_eq!(sink.payloads.len(), 2);
2119 assert_eq!(sink.payloads[0].row_token_offsets, [0]);
2120 assert_eq!(sink.payloads[1].row_token_offsets, [0]);
2121 }
2122
2123 #[test]
2124 fn write_direct_batch_to_sink_skips_send_for_empty_batch_but_records_stats() {
2125 let _trace_guard = direct_raw_trace_test_guard();
2126 let mappings = vec![mapping("id")];
2127 let mut state = WriterState::new(
2128 WriteBackend::DirectRawBulk,
2129 SchemaCheck::Strict,
2130 planned_schema(mappings),
2131 )
2132 .unwrap();
2133 let mut sink = RecordingRawSink::default();
2134 let batch = int32_batch("id", &[]);
2135
2136 let stats = poll_ready(write_direct_batch_to_sink(&mut state, &mut sink, &batch)).unwrap();
2137
2138 assert_eq!(
2139 stats,
2140 WriteStats {
2141 rows_written: 0,
2142 batches_written: 1
2143 }
2144 );
2145 assert!(sink.payloads.is_empty());
2146 }
2147
2148 #[test]
2149 fn write_direct_batch_to_sink_rejects_bad_later_row_before_send() {
2150 let _trace_guard = direct_raw_trace_test_guard();
2151 let mappings = vec![float_mapping("amount")];
2152 let mut state = WriterState::new(
2153 WriteBackend::DirectRawBulk,
2154 SchemaCheck::Strict,
2155 planned_schema(mappings),
2156 )
2157 .unwrap();
2158 let mut sink = RecordingRawSink::default();
2159 let batch = float64_batch("amount", &[Some(1.0), Some(f64::NAN)]);
2160
2161 let err =
2162 poll_ready(write_direct_batch_to_sink(&mut state, &mut sink, &batch)).unwrap_err();
2163
2164 assert_write_phase(&err, WritePhase::ValueConversion);
2165 let Error::ValueConversion { diagnostics } = inner_error(&err) else {
2166 panic!("expected value conversion error");
2167 };
2168 assert_eq!(diagnostics.all()[0].code(), DiagnosticCode::NonFiniteFloat);
2169 assert_eq!(diagnostics.all()[0].row(), Some(1));
2170 assert!(sink.payloads.is_empty());
2171 assert_eq!(state.stats(), WriteStats::default());
2172 }
2173
2174 #[test]
2175 fn write_direct_batch_to_sink_rejects_uint64_bigint_overflow_before_any_range_send() {
2176 let _trace_guard = direct_raw_trace_test_guard();
2177 let mappings = vec![schema_mapping_at(
2178 0,
2179 "u64_value",
2180 DataType::UInt64,
2181 MssqlType::BigInt,
2182 false,
2183 )];
2184 let mut state = WriterState::new(
2185 WriteBackend::DirectRawBulk,
2186 SchemaCheck::Strict,
2187 planned_schema(mappings),
2188 )
2189 .unwrap();
2190 let mut sink = RecordingRawSink::default();
2191 let row_count = DIRECT_RAW_MAX_PAYLOAD_BYTES / 9 + 2;
2192 let mut values = vec![1_u64; row_count];
2193 values[row_count - 1] = i64::MAX as u64 + 1;
2194 let batch = uint64_batch("u64_value", &values);
2195
2196 let err =
2197 poll_ready(write_direct_batch_to_sink(&mut state, &mut sink, &batch)).unwrap_err();
2198
2199 assert_write_phase(&err, WritePhase::ValueConversion);
2200 let Error::ValueConversion { diagnostics } = inner_error(&err) else {
2201 panic!("expected value conversion error");
2202 };
2203 assert_eq!(
2204 diagnostics.all()[0].code(),
2205 DiagnosticCode::IntegerOutOfRange
2206 );
2207 assert_eq!(diagnostics.all()[0].row(), Some(row_count - 1));
2208 assert!(sink.payloads.is_empty());
2209 assert_eq!(state.stats(), WriteStats::default());
2210 }
2211
2212 #[test]
2213 fn write_direct_batch_to_sink_rejects_runtime_type_mismatch_before_send() {
2214 let _trace_guard = direct_raw_trace_test_guard();
2215 let mappings = vec![mapping("id")];
2216 let mut state = WriterState::new(
2217 WriteBackend::DirectRawBulk,
2218 SchemaCheck::Strict,
2219 planned_schema(mappings),
2220 )
2221 .unwrap();
2222 let mut sink = RecordingRawSink::default();
2223 let batch = RecordBatch::try_new(
2224 Arc::new(Schema::new(vec![Field::new(
2225 "id",
2226 DataType::Float64,
2227 false,
2228 )])),
2229 vec![Arc::new(Float64Array::from(vec![1.0]))],
2230 )
2231 .unwrap();
2232
2233 let err =
2234 poll_ready(write_direct_batch_to_sink(&mut state, &mut sink, &batch)).unwrap_err();
2235
2236 assert_write_phase(&err, WritePhase::BatchSchemaValidation);
2237 let Error::ValueConversion { diagnostics } = inner_error(&err) else {
2238 panic!("expected value conversion error");
2239 };
2240 assert_eq!(diagnostics.all()[0].code(), DiagnosticCode::SchemaMismatch);
2241 assert!(
2242 diagnostics.all()[0]
2243 .message()
2244 .contains("runtime Arrow type Float64")
2245 );
2246 assert!(sink.payloads.is_empty());
2247 assert_eq!(state.stats(), WriteStats::default());
2248 }
2249
2250 #[test]
2251 fn write_direct_batch_to_sink_send_failure_preserves_error_and_keeps_stats() {
2252 let _trace_guard = direct_raw_trace_test_guard();
2253 let mappings = vec![mapping("id")];
2254 let mut state = WriterState::new(
2255 WriteBackend::DirectRawBulk,
2256 SchemaCheck::Strict,
2257 planned_schema(mappings),
2258 )
2259 .unwrap();
2260 let mut sink = RecordingRawSink {
2261 fail_on_send: true,
2262 payloads: Vec::new(),
2263 };
2264 let batch = int32_batch("id", &[1, 2, 3]);
2265
2266 let err =
2267 poll_ready(write_direct_batch_to_sink(&mut state, &mut sink, &batch)).unwrap_err();
2268
2269 assert_write_phase(&err, WritePhase::PacketWrite);
2270 let Error::Tiberius { source } = inner_error(&err) else {
2271 panic!("expected tiberius error");
2272 };
2273 assert_eq!(
2274 source.to_string(),
2275 "BULK UPLOAD input failure: fake raw send failure"
2276 );
2277 assert!(sink.payloads.is_empty());
2278 assert_eq!(state.stats(), WriteStats::default());
2279 }
2280
2281 #[test]
2282 fn writer_types_are_exported_from_crate_root() {
2283 assert_eq!(crate::WriteBackend::default(), WriteBackend::Auto);
2284 assert_eq!(crate::WriteOptions::default(), WriteOptions::default());
2285 assert_eq!(crate::WriteStats::default(), WriteStats::default());
2286 assert_eq!(crate::WritePhase::PacketWrite.as_str(), "packet_write");
2287 let _ = std::any::type_name::<crate::BulkWriter<'static, DummyStream>>();
2288 }
2289
2290 #[test]
2291 fn tiberius_alias_exposes_client_type() {
2292 let name = std::any::type_name::<tiberius::Client<DummyStream>>();
2293
2294 assert!(name.contains("tiberius"));
2295 }
2296
2297 fn assert_write_phase(error: &Error, expected: WritePhase) {
2298 assert_eq!(error.write_phase(), Some(expected));
2299 }
2300
2301 fn inner_error(error: &Error) -> &Error {
2302 error.without_write_phase()
2303 }
2304
2305 fn mapping(name: &str) -> SchemaMapping {
2306 SchemaMapping::new(
2307 ArrowFieldRef::new(0, name.to_owned(), false, DataType::Int32),
2308 MssqlColumn::new(Identifier::new(name).unwrap(), MssqlType::Int, false),
2309 )
2310 }
2311
2312 fn schema_mapping_at(
2313 index: usize,
2314 name: &str,
2315 arrow_type: DataType,
2316 mssql_type: MssqlType,
2317 nullable: bool,
2318 ) -> SchemaMapping {
2319 SchemaMapping::new(
2320 ArrowFieldRef::new(index, name.to_owned(), nullable, arrow_type),
2321 MssqlColumn::new(Identifier::new(name).unwrap(), mssql_type, nullable),
2322 )
2323 }
2324
2325 fn float_mapping(name: &str) -> SchemaMapping {
2326 float_mapping_at(0, name)
2327 }
2328
2329 fn float_mapping_at(index: usize, name: &str) -> SchemaMapping {
2330 SchemaMapping::new(
2331 ArrowFieldRef::new(index, name.to_owned(), false, DataType::Float64),
2332 MssqlColumn::new(
2333 Identifier::new(name).unwrap(),
2334 MssqlType::Float { precision: 53 },
2335 false,
2336 ),
2337 )
2338 }
2339
2340 fn utf8_mapping_at(index: usize, name: &str) -> SchemaMapping {
2341 SchemaMapping::new(
2342 ArrowFieldRef::new(index, name.to_owned(), false, DataType::Utf8),
2343 MssqlColumn::new(
2344 Identifier::new(name).unwrap(),
2345 MssqlType::NVarChar(MssqlTypeLength::Max),
2346 false,
2347 ),
2348 )
2349 }
2350
2351 fn binary_mapping_at(index: usize, name: &str) -> SchemaMapping {
2352 SchemaMapping::new(
2353 ArrowFieldRef::new(index, name.to_owned(), false, DataType::Binary),
2354 MssqlColumn::new(
2355 Identifier::new(name).unwrap(),
2356 MssqlType::VarBinary(MssqlTypeLength::Max),
2357 false,
2358 ),
2359 )
2360 }
2361
2362 fn fixed_size_binary_mapping_at(index: usize, name: &str, length: usize) -> SchemaMapping {
2363 SchemaMapping::new(
2364 ArrowFieldRef::new(
2365 index,
2366 name.to_owned(),
2367 false,
2368 DataType::FixedSizeBinary(i32::try_from(length).unwrap()),
2369 ),
2370 MssqlColumn::new(
2371 Identifier::new(name).unwrap(),
2372 MssqlType::Binary(length),
2373 false,
2374 ),
2375 )
2376 }
2377
2378 fn planned_schema(mappings: Vec<SchemaMapping>) -> PlannedSchema {
2379 planned_schema_with_options(mappings, PlanOptions::default())
2380 }
2381
2382 fn runtime_context_with_options(plan_options: PlanOptions) -> RuntimeConversionContext {
2383 RuntimeConversionContext::new(MssqlProfile::sql_server_2016_compat_100(), plan_options)
2384 }
2385
2386 fn planned_schema_with_options(
2387 mappings: Vec<SchemaMapping>,
2388 plan_options: PlanOptions,
2389 ) -> PlannedSchema {
2390 planned_schema_with_profile_and_options(
2391 MssqlProfile::sql_server_2016_compat_100(),
2392 plan_options,
2393 mappings,
2394 )
2395 }
2396
2397 fn planned_schema_with_profile_and_options(
2398 profile: MssqlProfile,
2399 plan_options: PlanOptions,
2400 mappings: Vec<SchemaMapping>,
2401 ) -> PlannedSchema {
2402 PlannedSchema::new(profile, plan_options, mappings)
2403 }
2404
2405 fn int32_batch(name: &str, values: &[i32]) -> RecordBatch {
2406 let schema = Arc::new(Schema::new(vec![Field::new(name, DataType::Int32, false)]));
2407 let array = Arc::new(Int32Array::from(values.to_vec()));
2408
2409 RecordBatch::try_new(schema, vec![array]).unwrap()
2410 }
2411
2412 fn uint64_batch(name: &str, values: &[u64]) -> RecordBatch {
2413 let schema = Arc::new(Schema::new(vec![Field::new(name, DataType::UInt64, false)]));
2414 let array = Arc::new(UInt64Array::from(values.to_vec()));
2415
2416 RecordBatch::try_new(schema, vec![array]).unwrap()
2417 }
2418
2419 fn binary_batch(name: &str, values: &[&[u8]]) -> RecordBatch {
2420 let schema = Arc::new(Schema::new(vec![Field::new(name, DataType::Binary, false)]));
2421 let array = Arc::new(BinaryArray::from_iter_values(values.iter().copied()));
2422
2423 RecordBatch::try_new(schema, vec![array]).unwrap()
2424 }
2425
2426 fn timestamp_microsecond_batch(name: &str, values: &[Option<i64>]) -> RecordBatch {
2427 let schema = Arc::new(Schema::new(vec![Field::new(
2428 name,
2429 DataType::Timestamp(TimeUnit::Microsecond, None),
2430 true,
2431 )]));
2432 let array = Arc::new(TimestampMicrosecondArray::from(values.to_vec()));
2433
2434 RecordBatch::try_new(schema, vec![array]).unwrap()
2435 }
2436
2437 fn bulk_target_column(ordinal: usize, name: &str, nullable: bool) -> FakeBulkTargetColumn {
2438 bulk_target_column_with_type(ordinal, name, nullable, tiberius::ColumnType::Int4)
2439 }
2440
2441 fn bulk_target_column_with_type(
2442 ordinal: usize,
2443 name: &str,
2444 nullable: bool,
2445 column_type: tiberius::ColumnType,
2446 ) -> FakeBulkTargetColumn {
2447 FakeBulkTargetColumn {
2448 ordinal,
2449 name: name.to_owned(),
2450 nullable,
2451 column_type,
2452 decimal_precision_scale: None,
2453 }
2454 }
2455
2456 fn bulk_target_decimal_column(
2457 ordinal: usize,
2458 name: &str,
2459 nullable: bool,
2460 precision: u8,
2461 scale: u8,
2462 ) -> FakeBulkTargetColumn {
2463 FakeBulkTargetColumn {
2464 ordinal,
2465 name: name.to_owned(),
2466 nullable,
2467 column_type: tiberius::ColumnType::Decimaln,
2468 decimal_precision_scale: Some((precision, scale)),
2469 }
2470 }
2471
2472 fn float64_batch(name: &str, values: &[Option<f64>]) -> RecordBatch {
2473 let schema = Arc::new(Schema::new(vec![Field::new(
2474 name,
2475 DataType::Float64,
2476 false,
2477 )]));
2478 let array = Arc::new(Float64Array::from(values.to_vec()));
2479
2480 RecordBatch::try_new(schema, vec![array]).unwrap()
2481 }
2482
2483 fn poll_ready<F>(future: F) -> F::Output
2484 where
2485 F: Future,
2486 {
2487 let mut context = Context::from_waker(Waker::noop());
2488 let mut future = Box::pin(future);
2489
2490 match future.as_mut().poll(&mut context) {
2491 Poll::Ready(output) => output,
2492 Poll::Pending => panic!("future unexpectedly returned pending"),
2493 }
2494 }
2495
2496 #[derive(Debug, Default)]
2497 struct RecordingSink {
2498 fail_on_send: Option<usize>,
2499 rows: Vec<tiberius::TokenRow<'static>>,
2500 }
2501
2502 #[derive(Debug, Default)]
2503 struct RecordingRawSink {
2504 fail_on_send: bool,
2505 payloads: Vec<RecordedRawPayload>,
2506 }
2507
2508 #[derive(Debug, PartialEq, Eq)]
2509 struct RecordedRawPayload {
2510 bytes: Vec<u8>,
2511 row_token_offsets: Vec<usize>,
2512 }
2513
2514 impl RawRowsSink for RecordingRawSink {
2515 async fn send_measured_raw_rows(
2516 &mut self,
2517 encoder: &DirectEncoder,
2518 batch: &RecordBatch,
2519 measured: &MeasuredDirectBatch,
2520 range: MeasuredRowRange,
2521 direct_observer: DirectRawBatchObserver,
2522 ) -> crate::Result<()> {
2523 let payload =
2524 encoder.encode_measured_batch_range(batch, measured, range.start, range.len)?;
2525
2526 if self.fail_on_send {
2527 return Err(Error::Tiberius {
2528 source: tiberius::error::Error::BulkInput(Cow::Borrowed(
2529 "fake raw send failure",
2530 )),
2531 });
2532 }
2533
2534 self.payloads.push(RecordedRawPayload {
2535 bytes: payload.bytes().to_vec(),
2536 row_token_offsets: payload.row_token_offsets().to_vec(),
2537 });
2538 direct_observer.record_packet_write_completed(
2539 range,
2540 payload.row_count(),
2541 payload.bytes().len(),
2542 std::time::Duration::ZERO,
2543 );
2544 Ok(())
2545 }
2546 }
2547
2548 impl TokenRowSink for RecordingSink {
2549 async fn send_token_row(&mut self, row: tiberius::TokenRow<'static>) -> crate::Result<()> {
2550 if self.fail_on_send == Some(self.rows.len()) {
2551 return Err(Error::Tiberius {
2552 source: tiberius::error::Error::BulkInput(Cow::Borrowed("fake send failure")),
2553 });
2554 }
2555
2556 self.rows.push(row);
2557 Ok(())
2558 }
2559 }
2560
2561 #[derive(Debug)]
2562 struct FakeBulkTargetColumn {
2563 ordinal: usize,
2564 name: String,
2565 nullable: bool,
2566 column_type: tiberius::ColumnType,
2567 decimal_precision_scale: Option<(u8, u8)>,
2568 }
2569
2570 impl BulkTargetColumnMetadata for FakeBulkTargetColumn {
2571 fn ordinal(&self) -> usize {
2572 self.ordinal
2573 }
2574
2575 fn name(&self) -> &str {
2576 &self.name
2577 }
2578
2579 fn is_nullable(&self) -> bool {
2580 self.nullable
2581 }
2582
2583 fn column_type(&self) -> tiberius::ColumnType {
2584 self.column_type
2585 }
2586
2587 fn decimal_precision_scale(&self) -> Option<(u8, u8)> {
2588 self.decimal_precision_scale
2589 }
2590 }
2591
2592 #[derive(Debug)]
2593 struct DummyStream;
2594
2595 impl AsyncRead for DummyStream {
2596 fn poll_read(
2597 self: Pin<&mut Self>,
2598 _cx: &mut Context<'_>,
2599 _buf: &mut [u8],
2600 ) -> Poll<std::io::Result<usize>> {
2601 Poll::Ready(Ok(0))
2602 }
2603 }
2604
2605 impl AsyncWrite for DummyStream {
2606 fn poll_write(
2607 self: Pin<&mut Self>,
2608 _cx: &mut Context<'_>,
2609 buf: &[u8],
2610 ) -> Poll<std::io::Result<usize>> {
2611 Poll::Ready(Ok(buf.len()))
2612 }
2613
2614 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2615 Poll::Ready(Ok(()))
2616 }
2617
2618 fn poll_close(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2619 Poll::Ready(Ok(()))
2620 }
2621 }
2622}