Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 6 additions & 8 deletions DbaClientX.SqlServer/SqlServer.BulkAutoCreate.cs
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,7 @@ private void EnsureAutoCreatedDestinationTable(
private void EnsureAutoCreatedDestinationTable(
SqlConnection connection,
SqlTransaction? transaction,
IDataReader reader,
IReadOnlyList<SqlServerBulkSourceColumn> columns,
string destinationTable,
SqlServerBulkInsertOptions? options)
{
Expand All @@ -65,7 +65,7 @@ private void EnsureAutoCreatedDestinationTable(
ExecuteBulkInsertSetupCommand(
connection,
transaction,
BuildCreateTableCommand(reader, destination, options.ColumnMappings),
BuildCreateTableCommand(columns, destination),
new Dictionary<string, object?> { ["@objectName"] = destination.QuotedFullName });
}

Expand Down Expand Up @@ -106,7 +106,7 @@ await ExecuteBulkInsertSetupCommandAsync(
private async Task EnsureAutoCreatedDestinationTableAsync(
SqlConnection connection,
SqlTransaction? transaction,
IDataReader reader,
IReadOnlyList<SqlServerBulkSourceColumn> columns,
string destinationTable,
SqlServerBulkInsertOptions? options,
CancellationToken cancellationToken)
Expand All @@ -131,7 +131,7 @@ await ExecuteBulkInsertSetupCommandAsync(
await ExecuteBulkInsertSetupCommandAsync(
connection,
transaction,
BuildCreateTableCommand(reader, destination, options.ColumnMappings),
BuildCreateTableCommand(columns, destination),
new Dictionary<string, object?> { ["@objectName"] = destination.QuotedFullName },
cancellationToken)
.ConfigureAwait(false);
Expand Down Expand Up @@ -229,11 +229,9 @@ private static string BuildCreateTableCommand(
}

private static string BuildCreateTableCommand(
IDataReader reader,
SqlServerDestinationTable destination,
IDictionary<string, string>? columnMappings)
IReadOnlyList<SqlServerBulkSourceColumn> columns,
SqlServerDestinationTable destination)
{
var columns = GetReaderColumns(reader, columnMappings);
var builder = new StringBuilder();
builder.AppendLine("IF OBJECT_ID(@objectName, N'U') IS NULL");
builder.AppendLine("BEGIN");
Expand Down
63 changes: 31 additions & 32 deletions DbaClientX.SqlServer/SqlServer.BulkDataReader.cs
Original file line number Diff line number Diff line change
Expand Up @@ -25,8 +25,6 @@ public virtual void BulkInsert(
string? username = null,
string? password = null)
{
ValidateBulkInsertInputs(reader, destinationTable, batchSize, bulkCopyTimeout, options);

var connectionString = BuildConnectionString(serverOrInstance, database, integratedSecurity, username, password);
BulkInsert(connectionString, reader, destinationTable, options, useTransaction, batchSize, bulkCopyTimeout);
}
Expand Down Expand Up @@ -60,7 +58,7 @@ public virtual void BulkInsert(
int? bulkCopyTimeout = null)
{
ValidateConnectionString(connectionString);
ValidateBulkInsertInputs(reader, destinationTable, batchSize, bulkCopyTimeout, options);
var columns = ValidateBulkInsertInputs(reader, destinationTable, batchSize, bulkCopyTimeout, options);

SqlConnection? connection = null;
SqlTransaction? transaction = null;
Expand All @@ -69,9 +67,9 @@ public virtual void BulkInsert(
try
{
(connection, transaction, dispose) = ResolveConnection(connectionString, useTransaction);
EnsureAutoCreatedDestinationTable(connection!, transaction, reader, destinationTable, options);
EnsureAutoCreatedDestinationTable(connection!, transaction, columns, destinationTable, options);
using var bulkCopy = CreateBulkCopy(connection!, transaction, options);
ConfigureBulkCopy(bulkCopy, reader, destinationTable, batchSize, bulkCopyTimeout, options);
ConfigureBulkCopy(bulkCopy, columns, destinationTable, batchSize, bulkCopyTimeout, options);
WriteToServer(bulkCopy, reader);
}
catch (DbaTransactionException)
Expand Down Expand Up @@ -120,8 +118,6 @@ public virtual async Task BulkInsertAsync(
string? username = null,
string? password = null)
{
ValidateBulkInsertInputs(reader, destinationTable, batchSize, bulkCopyTimeout, options);

var connectionString = BuildConnectionString(serverOrInstance, database, integratedSecurity, username, password);
await BulkInsertAsync(connectionString, reader, destinationTable, options, useTransaction, batchSize, bulkCopyTimeout, cancellationToken).ConfigureAwait(false);
}
Expand Down Expand Up @@ -157,7 +153,7 @@ public virtual async Task BulkInsertAsync(
CancellationToken cancellationToken = default)
{
ValidateConnectionString(connectionString);
ValidateBulkInsertInputs(reader, destinationTable, batchSize, bulkCopyTimeout, options);
var columns = ValidateBulkInsertInputs(reader, destinationTable, batchSize, bulkCopyTimeout, options);

SqlConnection? connection = null;
SqlTransaction? transaction = null;
Expand All @@ -166,9 +162,9 @@ public virtual async Task BulkInsertAsync(
try
{
(connection, transaction, dispose) = await ResolveConnectionAsync(connectionString, useTransaction, cancellationToken).ConfigureAwait(false);
await EnsureAutoCreatedDestinationTableAsync(connection!, transaction, reader, destinationTable, options, cancellationToken).ConfigureAwait(false);
await EnsureAutoCreatedDestinationTableAsync(connection!, transaction, columns, destinationTable, options, cancellationToken).ConfigureAwait(false);
using var bulkCopy = CreateBulkCopy(connection!, transaction, options);
ConfigureBulkCopy(bulkCopy, reader, destinationTable, batchSize, bulkCopyTimeout, options);
ConfigureBulkCopy(bulkCopy, columns, destinationTable, batchSize, bulkCopyTimeout, options);
await AwaitWithCallerCancellationAsync(
() => WriteToServerAsync(bulkCopy, reader, cancellationToken),
cancellationToken).ConfigureAwait(false);
Expand Down Expand Up @@ -200,7 +196,7 @@ public virtual Task BulkInsertAsync(
CancellationToken cancellationToken = default)
=> BulkInsertAsync(connectionString, reader, destinationTable, options: null, useTransaction, batchSize, bulkCopyTimeout, cancellationToken);

private static void ConfigureBulkCopy(SqlBulkCopy bulkCopy, IDataReader reader, string destinationTable, int? batchSize, int? bulkCopyTimeout, SqlServerBulkInsertOptions? options)
private static void ConfigureBulkCopy(SqlBulkCopy bulkCopy, IReadOnlyList<SqlServerBulkSourceColumn> columns, string destinationTable, int? batchSize, int? bulkCopyTimeout, SqlServerBulkInsertOptions? options)
{
bulkCopy.DestinationTableName = ResolveBulkCopyDestinationTableName(destinationTable, options);
bulkCopy.EnableStreaming = true;
Expand All @@ -226,13 +222,13 @@ private static void ConfigureBulkCopy(SqlBulkCopy bulkCopy, IDataReader reader,
bulkCopy.NotifyAfter = notifyAfter.Value;
}

foreach (var column in GetReaderColumns(reader, options?.ColumnMappings))
foreach (var column in columns)
{
bulkCopy.ColumnMappings.Add(column.Ordinal, column.DestinationName);
}
}

private static void ValidateBulkInsertInputs(IDataReader reader, string destinationTable, int? batchSize, int? bulkCopyTimeout, SqlServerBulkInsertOptions? options = null)
private static IReadOnlyList<SqlServerBulkSourceColumn> ValidateBulkInsertInputs(IDataReader reader, string destinationTable, int? batchSize, int? bulkCopyTimeout, SqlServerBulkInsertOptions? options = null)
{
if (reader == null)
{
Expand Down Expand Up @@ -264,10 +260,10 @@ private static void ValidateBulkInsertInputs(IDataReader reader, string destinat
throw new ArgumentOutOfRangeException(nameof(SqlServerBulkInsertOptions.NotifyAfter), "NotifyAfter must be greater than zero.");
}

ValidateColumnMappings(reader, options?.ColumnMappings);
return GetValidatedReaderColumns(reader, options?.ColumnMappings);
}

private static void ValidateColumnMappings(IDataReader reader, IDictionary<string, string>? columnMappings)
private static IReadOnlyList<SqlServerBulkSourceColumn> GetValidatedReaderColumns(IDataReader reader, IDictionary<string, string>? columnMappings)
{
if (columnMappings != null)
{
Expand All @@ -283,35 +279,38 @@ private static void ValidateColumnMappings(IDataReader reader, IDictionary<strin
throw new ArgumentException("Column mapping destination cannot be null or whitespace.", nameof(columnMappings));
}

if (!ContainsColumn(reader, mapping.Key, columnMappings))
}
}

var columns = GetReaderColumns(reader, columnMappings);
if (columnMappings != null)
{
var comparer = GetComparer(columnMappings);
var sourceColumns = new HashSet<string>(comparer);
foreach (var column in columns)
{
sourceColumns.Add(column.SourceName);
}

foreach (var mapping in columnMappings)
{
if (!sourceColumns.Contains(mapping.Key))
{
throw new ArgumentException($"Column mapping source '{mapping.Key}' does not exist in the source reader.", nameof(columnMappings));
}
}
}

var destinationColumns = new HashSet<string>(StringComparer.OrdinalIgnoreCase);
foreach (var column in GetReaderColumns(reader, columnMappings))
foreach (var column in columns)
{
if (!destinationColumns.Add(column.DestinationName))
{
throw new ArgumentException($"Column mappings produce duplicate destination column '{column.DestinationName}'.", nameof(columnMappings));
}
}
}

private static bool ContainsColumn(IDataReader reader, string columnName, IDictionary<string, string> columnMappings)
{
var comparer = GetComparer(columnMappings);
foreach (var sourceName in GetReaderSourceNames(reader, columnMappings))
{
if (comparer.Equals(sourceName, columnName))
{
return true;
}
}

return false;
return columns;
}

private static IReadOnlyList<SqlServerBulkSourceColumn> GetReaderColumns(IDataReader reader, IDictionary<string, string>? columnMappings)
Expand Down Expand Up @@ -370,13 +369,13 @@ private static string GetUniqueReaderSourceName(IDataReader reader, int ordinal,

private static Dictionary<int, DataRow> GetReaderSchemaRows(IDataReader reader)
{
var rows = new Dictionary<int, DataRow>();
var schema = reader.GetSchemaTable();
if (schema == null)
{
return rows;
return new Dictionary<int, DataRow>();
}

var rows = new Dictionary<int, DataRow>(schema.Rows.Count);
foreach (DataRow row in schema.Rows)
{
if (TryGetSchemaValue<int>(row, "ColumnOrdinal", out var ordinal))
Expand Down
5 changes: 5 additions & 0 deletions DbaClientX.Tests/SqlServerBulkInsertTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,8 @@ public UnnamedColumnReader(string[] names)

public int FieldCount => _names.Length;

public int SchemaTableCallCount { get; private set; }

public void Close() => IsClosed = true;

public void Dispose() => Close();
Expand Down Expand Up @@ -219,6 +221,7 @@ public UnnamedColumnReader(string[] names)

public DataTable GetSchemaTable()
{
SchemaTableCallCount++;
var schema = new DataTable();
schema.Columns.Add("ColumnName", typeof(string));
schema.Columns.Add("ColumnOrdinal", typeof(int));
Expand Down Expand Up @@ -736,6 +739,7 @@ public void BulkInsert_WithUnnamedReaderColumn_AllowsMappingSynthesizedSourceNam

sqlServer.BulkInsert("s", "db", true, reader, "dbo.Dest", options);

Assert.Equal(1, reader.SchemaTableCallCount);
Assert.Contains(sqlServer.Mappings, mapping => mapping.Source == "0" && mapping.Destination == "Id");
}

Expand All @@ -755,6 +759,7 @@ public void BulkInsert_WithUnnamedReaderColumn_KeepsSynthesizedSourceNamesUnique

sqlServer.BulkInsert("s", "db", true, reader, "dbo.Dest", options);

Assert.Equal(1, reader.SchemaTableCallCount);
Assert.Contains(sqlServer.Mappings, mapping => mapping.Source == "0" && mapping.Destination == "GeneratedId");
Assert.Contains(sqlServer.Mappings, mapping => mapping.Source == "1" && mapping.Destination == "ExistingColumn");
}
Expand Down
Loading