Skip to main content

arrow_sql_server/write/
writer.rs

1//! Baseline bulk writer public API skeleton.
2
3use 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/// Write backend selection.
36#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
37pub enum WriteBackend {
38    /// Select the best available backend for the current crate build and plan.
39    #[default]
40    Auto,
41    /// Use Tiberius' row-oriented `TokenRow` bulk-load path.
42    BaselineTokenRow,
43    /// Use direct bulk-row payload encoding through Tiberius' framed sink.
44    DirectFramedBulk,
45    /// Use the raw bulk-row payload path exposed by the Tiberius fork.
46    DirectRawBulk,
47}
48
49/// Execution-time write options.
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
51pub struct WriteOptions {
52    /// Requested write backend.
53    pub backend: WriteBackend,
54    /// Batch schema validation policy.
55    pub schema_check: SchemaCheck,
56}
57
58/// Cumulative write statistics.
59#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Hash)]
60pub struct WriteStats {
61    /// Number of rows accepted by the writer.
62    pub rows_written: u64,
63    /// Number of batches accepted by the writer.
64    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/// SQL Server bulk writer for Arrow record batches.
136#[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    /// Starts a bulk writer for a planned SQL Server table target.
150    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    /// Writes one Arrow record batch.
248    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    /// Finalizes the bulk writer and returns cumulative write statistics.
261    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
695/// Writes one baseline token-row batch without crate-owned batch lifecycle traces.
696async 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
732/// Adds the crate-owned batch lifecycle span around the baseline token-row write path.
733async 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/// Test-only direct write entry point with direct raw detail telemetry disabled.
763#[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
776/// Adds the crate-owned batch lifecycle span and direct raw detail telemetry.
777async 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
797/// Shared direct write implementation.
798///
799/// The `direct_observer` records direct raw detail events when enabled. Batch
800/// lifecycle tracing is owned by `write_traced_direct_batch_to_sink`.
801async 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}