perf(import): bound bulk import memory and cancellation

This commit is contained in:
miracle
2026-07-29 22:53:34 +08:00
committed by GitHub
parent 843f945581
commit d0c3a3fd2e
7 changed files with 3761 additions and 193 deletions
@@ -0,0 +1,550 @@
use std::fs::{self, File};
use std::io::{BufWriter, Write};
use std::path::Path;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::{Duration, Instant};
use dbx_core::connection::AppState;
use dbx_core::models::connection::{ConnectionConfig, DatabaseType};
use dbx_core::storage::Storage;
use dbx_core::table_import::{
import_table_file_core, TableImportColumnMapping, TableImportMode, TableImportParseOptions, TableImportPhase,
TableImportRequest, TableImportSourceFormat,
};
use dbx_core::xlsx_export::{build_xlsx_workbook, XlsxWorksheetData};
use serde_json::json;
use sysinfo::{get_current_pid, ProcessRefreshKind, ProcessesToUpdate, RefreshKind, System};
#[derive(Debug, Clone, Copy)]
enum BenchDatabase {
Mysql,
Postgres,
SqlServer,
}
impl BenchDatabase {
fn parse(value: &str) -> Result<Self, String> {
match value {
"mysql" => Ok(Self::Mysql),
"postgres" => Ok(Self::Postgres),
"sqlserver" => Ok(Self::SqlServer),
_ => Err(format!("Unsupported database: {value}")),
}
}
fn db_type(self) -> DatabaseType {
match self {
Self::Mysql => DatabaseType::Mysql,
Self::Postgres => DatabaseType::Postgres,
Self::SqlServer => DatabaseType::SqlServer,
}
}
fn label(self) -> &'static str {
match self {
Self::Mysql => "mysql",
Self::Postgres => "postgres",
Self::SqlServer => "sqlserver",
}
}
fn default_port(self) -> u16 {
match self {
Self::Mysql => 3306,
Self::Postgres => 5432,
Self::SqlServer => 1433,
}
}
}
#[derive(Debug, Clone, Copy)]
enum BenchFormat {
Csv,
Xlsx,
}
impl BenchFormat {
fn parse(value: &str) -> Result<Self, String> {
match value {
"csv" => Ok(Self::Csv),
"xlsx" => Ok(Self::Xlsx),
_ => Err(format!("Unsupported format: {value}")),
}
}
fn source_format(self) -> TableImportSourceFormat {
match self {
Self::Csv => TableImportSourceFormat::Csv,
Self::Xlsx => TableImportSourceFormat::Excel,
}
}
fn extension(self) -> &'static str {
match self {
Self::Csv => "csv",
Self::Xlsx => "xlsx",
}
}
}
struct Options {
database: BenchDatabase,
format: BenchFormat,
rows: usize,
columns: usize,
batch_size: usize,
text_bytes: usize,
}
fn print_help() {
println!(
"Live table import benchmark\n\n\
Usage:\n cargo run -p dbx-core --example table_import_live_bench --release -- [options]\n\n\
Options:\n\
--database=postgres mysql, postgres, or sqlserver\n\
--format=csv csv or xlsx\n\
--rows=200000 Number of data rows\n\
--columns=12 Number of columns\n\
--batch-size=500 Import batch size\n\
--text-bytes=0 Exact text value bytes; 0 keeps the default values\n\n\
Connection environment:\n\
DBX_BENCH_HOST, DBX_BENCH_PORT, DBX_BENCH_USER, DBX_BENCH_PASSWORD,\n\
DBX_BENCH_DATABASE, DBX_BENCH_SCHEMA, DBX_BENCH_SSL"
);
}
fn parse_options() -> Result<Options, String> {
let mut database = BenchDatabase::Postgres;
let mut format = BenchFormat::Csv;
let mut rows = 200_000;
let mut columns = 12;
let mut batch_size = 500;
let mut text_bytes = 0;
for argument in std::env::args().skip(1) {
if matches!(argument.as_str(), "--help" | "-h") {
print_help();
std::process::exit(0);
}
let (key, value) = argument.split_once('=').ok_or_else(|| format!("Invalid option: {argument}"))?;
match key {
"--database" => database = BenchDatabase::parse(value)?,
"--format" => format = BenchFormat::parse(value)?,
"--rows" => rows = value.parse().map_err(|_| format!("Invalid row count: {value}"))?,
"--columns" => columns = value.parse().map_err(|_| format!("Invalid column count: {value}"))?,
"--batch-size" => batch_size = value.parse().map_err(|_| format!("Invalid batch size: {value}"))?,
"--text-bytes" => text_bytes = value.parse().map_err(|_| format!("Invalid text byte count: {value}"))?,
_ => return Err(format!("Unknown option: {key}")),
}
}
if rows == 0 || columns < 2 || batch_size == 0 {
return Err("rows and batch-size must be positive; columns must be at least 2".to_string());
}
Ok(Options { database, format, rows, columns, batch_size, text_bytes })
}
fn env_required(name: &str) -> Result<String, String> {
std::env::var(name).map_err(|_| format!("Missing environment variable: {name}"))
}
fn connection_config(id: &str, database: BenchDatabase) -> Result<ConnectionConfig, String> {
let database_name = env_required("DBX_BENCH_DATABASE")?;
Ok(ConnectionConfig {
id: id.to_string(),
name: id.to_string(),
note: String::new(),
db_type: database.db_type(),
driver_profile: None,
driver_label: None,
url_params: None,
agent_java_options: Vec::new(),
host: env_required("DBX_BENCH_HOST")?,
port: std::env::var("DBX_BENCH_PORT")
.ok()
.and_then(|value| value.parse().ok())
.unwrap_or_else(|| database.default_port()),
username: env_required("DBX_BENCH_USER")?,
password: env_required("DBX_BENCH_PASSWORD")?,
database: Some(database_name),
visible_databases: None,
visible_schemas: None,
show_system_schemas: false,
attached_databases: Vec::new(),
init_script: None,
color: None,
transport_layers: Vec::new(),
connect_timeout_secs: 15,
query_timeout_secs: 120,
idle_timeout_secs: 60,
keepalive_interval_secs: 0,
ssl: std::env::var("DBX_BENCH_SSL").is_ok_and(|value| value.eq_ignore_ascii_case("true")),
ca_cert_path: String::new(),
client_cert_path: String::new(),
client_key_path: String::new(),
sysdba: false,
oracle_connection_type: None,
connection_string: None,
redis_connection_mode: None,
redis_sentinel_master: String::new(),
redis_sentinel_nodes: String::new(),
redis_sentinel_username: String::new(),
redis_sentinel_password: String::new(),
redis_sentinel_tls: false,
redis_cluster_nodes: String::new(),
redis_key_separator: dbx_core::models::connection::default_redis_key_separator(),
redis_scan_page_size: None,
redis_database_aliases: Default::default(),
etcd_endpoints: String::new(),
gbase_server: String::new(),
informix_server: String::new(),
external_config: None,
jdbc_driver_class: None,
jdbc_driver_paths: Vec::new(),
one_time: false,
read_only: false,
is_production: false,
production_databases: Vec::new(),
database_info: None,
})
}
fn columns(count: usize) -> Vec<String> {
(0..count).map(|index| format!("column_{}", index + 1)).collect()
}
fn row(row_index: usize, column_count: usize, text_bytes: usize) -> Vec<serde_json::Value> {
(0..column_count)
.map(|column_index| {
if column_index == 0 {
json!(row_index + 1)
} else {
let value = format!("value-{row_index:08}-{column_index:02}");
if text_bytes == 0 {
json!(value)
} else if value.len() >= text_bytes {
json!(&value[..text_bytes])
} else {
json!(format!("{value}{}", "x".repeat(text_bytes - value.len())))
}
}
})
.collect()
}
fn write_csv(path: &Path, row_count: usize, column_count: usize, text_bytes: usize) -> Result<(), String> {
let mut writer = BufWriter::new(File::create(path).map_err(|error| error.to_string())?);
writeln!(writer, "{}", columns(column_count).join(",")).map_err(|error| error.to_string())?;
for row_index in 0..row_count {
writeln!(
writer,
"{}",
row(row_index, column_count, text_bytes)
.into_iter()
.map(|value| value.as_str().map(str::to_string).unwrap_or_else(|| value.to_string()))
.collect::<Vec<_>>()
.join(",")
)
.map_err(|error| error.to_string())?;
}
writer.flush().map_err(|error| error.to_string())
}
fn write_xlsx(path: &Path, row_count: usize, column_count: usize, text_bytes: usize) -> Result<(), String> {
let workbook = build_xlsx_workbook(&XlsxWorksheetData {
sheet_name: Some("Benchmark".to_string()),
columns: columns(column_count),
column_types: Vec::new(),
rows: (0..row_count).map(|row_index| row(row_index, column_count, text_bytes)).collect(),
numeric_column_right_align: false,
})?;
fs::write(path, workbook).map_err(|error| error.to_string())
}
struct PeakRssSampler {
baseline_bytes: u64,
peak_bytes: Arc<AtomicU64>,
stop: Arc<AtomicBool>,
thread: Option<thread::JoinHandle<()>>,
}
impl PeakRssSampler {
fn start() -> Result<Self, String> {
let pid = get_current_pid().map_err(|error| error.to_string())?;
let mut system =
System::new_with_specifics(RefreshKind::new().with_processes(ProcessRefreshKind::new().with_memory()));
system.refresh_processes_specifics(
ProcessesToUpdate::Some(&[pid]),
true,
ProcessRefreshKind::new().with_memory(),
);
let baseline_bytes = system.process(pid).map(|process| process.memory()).unwrap_or(0);
let peak_bytes = Arc::new(AtomicU64::new(baseline_bytes));
let stop = Arc::new(AtomicBool::new(false));
let peak_for_thread = peak_bytes.clone();
let stop_for_thread = stop.clone();
let handle = thread::spawn(move || {
let mut system =
System::new_with_specifics(RefreshKind::new().with_processes(ProcessRefreshKind::new().with_memory()));
while !stop_for_thread.load(Ordering::Relaxed) {
system.refresh_processes_specifics(
ProcessesToUpdate::Some(&[pid]),
true,
ProcessRefreshKind::new().with_memory(),
);
if let Some(process) = system.process(pid) {
peak_for_thread.fetch_max(process.memory(), Ordering::Relaxed);
}
thread::sleep(Duration::from_millis(10));
}
});
Ok(Self { baseline_bytes, peak_bytes, stop, thread: Some(handle) })
}
fn stop(mut self) -> (u64, u64) {
self.shutdown();
(self.baseline_bytes, self.peak_bytes.load(Ordering::Relaxed))
}
fn shutdown(&mut self) {
self.stop.store(true, Ordering::Relaxed);
if let Some(handle) = self.thread.take() {
let _ = handle.join();
}
}
}
impl Drop for PeakRssSampler {
fn drop(&mut self) {
self.shutdown();
}
}
fn qualified_table(database: BenchDatabase, schema: &str, table: &str) -> String {
match database {
BenchDatabase::Mysql => format!("`{schema}`.`{table}`"),
BenchDatabase::Postgres => format!("\"{schema}\".\"{table}\""),
BenchDatabase::SqlServer => format!("[{schema}].[{table}]"),
}
}
async fn execute_sql(
state: &AppState,
connection_id: &str,
database: &str,
schema: &str,
sql: &str,
) -> Result<(), String> {
dbx_core::query::execute_sql_statement(state, connection_id, database, sql, Some(schema), None).await.map(|_| ())
}
fn create_table_sql(database: BenchDatabase, schema: &str, table: &str, column_count: usize) -> Vec<String> {
let qualified = qualified_table(database, schema, table);
let definitions = columns(column_count)
.into_iter()
.enumerate()
.map(|(index, column)| match database {
BenchDatabase::Mysql if index == 0 => format!("`{column}` BIGINT NOT NULL"),
BenchDatabase::Mysql => format!("`{column}` TEXT NULL"),
BenchDatabase::Postgres if index == 0 => format!("\"{column}\" BIGINT NOT NULL"),
BenchDatabase::Postgres => format!("\"{column}\" TEXT NULL"),
BenchDatabase::SqlServer if index == 0 => format!("[{column}] BIGINT NOT NULL"),
BenchDatabase::SqlServer => format!("[{column}] NVARCHAR(200) NULL"),
})
.collect::<Vec<_>>()
.join(", ");
let mut statements = Vec::new();
if matches!(database, BenchDatabase::Postgres) {
statements.push(format!("CREATE SCHEMA IF NOT EXISTS \"{schema}\""));
}
statements.push(format!("DROP TABLE IF EXISTS {qualified}"));
statements.push(format!("CREATE TABLE {qualified} ({definitions})"));
statements
}
fn import_request(
connection_id: &str,
database: &str,
schema: &str,
table: &str,
path: &Path,
options: &Options,
import_id: &str,
) -> TableImportRequest {
TableImportRequest {
import_id: import_id.to_string(),
connection_id: connection_id.to_string(),
database: database.to_string(),
schema: schema.to_string(),
table: table.to_string(),
file_path: path.to_string_lossy().to_string(),
source_ref: None,
source_format: Some(options.format.source_format()),
parse_options: TableImportParseOptions::default(),
mappings: columns(options.columns)
.into_iter()
.map(|column| TableImportColumnMapping {
source_column: column.clone(),
target_column: column,
target_data_type: None,
})
.collect(),
mode: TableImportMode::Append,
create_table: false,
batch_size: options.batch_size,
date_time_format: None,
prepared_source: None,
retain_source: false,
}
}
async fn run() -> Result<(), String> {
let options = parse_options()?;
let database_name = env_required("DBX_BENCH_DATABASE")?;
let schema = env_required("DBX_BENCH_SCHEMA")?;
let suffix = uuid::Uuid::new_v4().simple().to_string();
let connection_id = format!("table-import-live-bench-{suffix}");
let table = format!("dbx_import_bench_{}", &suffix[..12]);
let temp_dir = tempfile::Builder::new()
.prefix(&format!("dbx-table-import-live-bench-{suffix}-"))
.tempdir()
.map_err(|error| error.to_string())?;
let source_path = temp_dir.path().join(format!("benchmark.{}", options.format.extension()));
match options.format {
BenchFormat::Csv => write_csv(&source_path, options.rows, options.columns, options.text_bytes)?,
BenchFormat::Xlsx => write_xlsx(&source_path, options.rows, options.columns, options.text_bytes)?,
}
let file_bytes = fs::metadata(&source_path).map_err(|error| error.to_string())?.len();
let storage = Storage::open(&temp_dir.path().join("storage.db")).await?;
let state = AppState::new(storage);
let config = connection_config(&connection_id, options.database)?;
state.configs.write().await.insert(connection_id.clone(), config);
let pool_key = state.get_or_create_pool(&connection_id, Some(&database_name)).await?;
let benchmark_result: Result<serde_json::Value, String> = async {
for sql in create_table_sql(options.database, &schema, &table, options.columns) {
execute_sql(&state, &connection_id, &database_name, &schema, &sql).await?;
}
let throughput_request = import_request(
&connection_id,
&database_name,
&schema,
&table,
&source_path,
&options,
&format!("throughput-{suffix}"),
);
let rss = PeakRssSampler::start()?;
let throughput_started = Instant::now();
let throughput_summary = import_table_file_core(
&state,
&throughput_request,
&options.database.db_type(),
&pool_key,
|_| Box::pin(async { false }),
|_| {},
)
.await?;
let throughput_elapsed = throughput_started.elapsed();
let (baseline_rss_bytes, peak_rss_bytes) = rss.stop();
if throughput_summary.rows_imported != options.rows {
return Err(format!(
"Import row count mismatch: expected {}, imported {}",
options.rows, throughput_summary.rows_imported
));
}
execute_sql(
&state,
&connection_id,
&database_name,
&schema,
&format!("TRUNCATE TABLE {}", qualified_table(options.database, &schema, &table)),
)
.await?;
let cancel_requested = Arc::new(AtomicBool::new(false));
let cancel_requested_at = Arc::new(Mutex::new(None::<Instant>));
let cancel_for_check = cancel_requested.clone();
let cancel_for_progress = cancel_requested.clone();
let cancel_time_for_progress = cancel_requested_at.clone();
let cancel_request = import_request(
&connection_id,
&database_name,
&schema,
&table,
&source_path,
&options,
&format!("cancel-{suffix}"),
);
let cancel_result = import_table_file_core(
&state,
&cancel_request,
&options.database.db_type(),
&pool_key,
move |_| {
let cancel = cancel_for_check.clone();
Box::pin(async move { cancel.load(Ordering::Acquire) })
},
move |progress| {
if progress.phase == TableImportPhase::Writing
&& progress.rows_imported > 0
&& cancel_for_progress.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire).is_ok()
{
*cancel_time_for_progress.lock().expect("cancel time lock") = Some(Instant::now());
}
},
)
.await;
let cancel_latency = cancel_requested_at
.lock()
.map_err(|_| "cancel time lock poisoned".to_string())?
.as_ref()
.map(Instant::elapsed)
.ok_or_else(|| "import completed before cancellation was requested".to_string())?;
if !cancel_result.as_ref().is_err_and(|error| error.contains("cancelled")) {
return Err(format!("Expected cancelled import, got: {cancel_result:?}"));
}
let rows_imported = throughput_summary.rows_imported;
let elapsed_seconds = throughput_elapsed.as_secs_f64();
Ok(json!({
"database": options.database.label(),
"format": options.format.extension(),
"fileBytes": file_bytes,
"rows": options.rows,
"columns": options.columns,
"batchSize": options.batch_size,
"textBytes": options.text_bytes,
"rowsImported": rows_imported,
"elapsedMs": throughput_elapsed.as_secs_f64() * 1000.0,
"rowsPerSecond": rows_imported as f64 / elapsed_seconds,
"baselineRssBytes": baseline_rss_bytes,
"peakRssBytes": peak_rss_bytes,
"peakRssDeltaBytes": peak_rss_bytes.saturating_sub(baseline_rss_bytes),
"cancellationLatencyMs": cancel_latency.as_secs_f64() * 1000.0,
}))
}
.await;
let cleanup_result = execute_sql(
&state,
&connection_id,
&database_name,
&schema,
&format!("DROP TABLE IF EXISTS {}", qualified_table(options.database, &schema, &table)),
)
.await;
state.remove_connection_pools_detached(&connection_id).await;
let output = benchmark_result?;
cleanup_result?;
println!("{}", serde_json::to_string_pretty(&output).map_err(|error| error.to_string())?);
Ok(())
}
#[tokio::main]
async fn main() {
if let Err(error) = run().await {
eprintln!("{error}");
std::process::exit(1);
}
}
+29
View File
@@ -25,6 +25,7 @@ use super::file_validator::validate_file_path;
pub type MySqlPool = mysql_async::Pool;
const MYSQL_TCP_KEEPALIVE_MS: u32 = 30_000;
const MYSQL_SQL_PACKET_MARGIN_MAX_BYTES: usize = 64 * 1024;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum MySqlCatalogDialect {
@@ -3655,6 +3656,23 @@ pub async fn execute_query(pool: &MySqlPool, sql: &str, bare: bool) -> Result<Qu
execute_query_with_max_rows(pool, sql, bare, None, MySqlQueryDialect::default()).await
}
pub async fn max_allowed_packet(pool: &MySqlPool) -> Result<u64, String> {
let mut conn = get_conn_with_health_check(pool).await?;
conn.query_first::<u64, _>("SELECT @@max_allowed_packet")
.await
.map_err(|error| error.to_string())?
.ok_or_else(|| "MySQL did not return @@max_allowed_packet".to_string())
}
pub(crate) fn mysql_sql_statement_hard_limit(max_allowed_packet: u64) -> Option<usize> {
let packet_bytes = usize::try_from(max_allowed_packet).ok()?;
if packet_bytes == 0 {
return None;
}
let margin = (packet_bytes / 10).clamp(1024, MYSQL_SQL_PACKET_MARGIN_MAX_BYTES).min(packet_bytes / 2);
packet_bytes.checked_sub(margin).filter(|limit| *limit > 0)
}
pub async fn execute_query_with_max_rows(
pool: &MySqlPool,
sql: &str,
@@ -4158,6 +4176,17 @@ pub async fn list_triggers(pool: &MySqlPool, database: &str, table: &str) -> Res
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mysql_sql_statement_limit_reserves_packet_headroom() {
let packet_bytes = 64 * 1024 * 1024;
let hard_limit = mysql_sql_statement_hard_limit(packet_bytes).unwrap();
assert!(hard_limit < packet_bytes as usize);
assert!(hard_limit >= packet_bytes as usize * 9 / 10);
assert_eq!(mysql_sql_statement_hard_limit(0), None);
assert_eq!(mysql_sql_statement_hard_limit(4096), Some(3072));
}
use crate::db::connection_timeout;
use mysql_async::consts::ColumnFlags;
#[test]
+95 -2
View File
@@ -8,11 +8,14 @@ use futures::{FutureExt, TryStreamExt};
use sqlparser::ast::{Expr, Ident, SelectItem, SetExpr, Statement};
use sqlparser::dialect::MsSqlDialect;
use sqlparser::parser::Parser;
use std::borrow::Cow;
use std::future::Future;
use std::panic::AssertUnwindSafe;
use std::sync::{Arc as StdArc, Mutex as StdMutex};
use std::time::{Duration, Instant};
use tiberius::{AuthMethod, Client, ColumnData, ColumnType, Config, FromSql, QueryItem, QueryStream, Row, SqlBrowser};
use tiberius::{
AuthMethod, Client, ColumnData, ColumnType, Config, FromSql, QueryItem, QueryStream, Row, SqlBrowser, TokenRow,
};
use tokio::net::TcpStream;
use tokio_util::compat::{Compat, TokioAsyncWriteCompatExt};
use tokio_util::sync::CancellationToken;
@@ -2084,6 +2087,57 @@ pub async fn execute_query(client: &mut SqlServerClient, sql: &str) -> Result<Qu
execute_query_with_max_rows(client, sql, None).await
}
fn sqlserver_bulk_token_row(values: Vec<Option<String>>) -> TokenRow<'static> {
let mut row = TokenRow::with_capacity(values.len());
for value in values {
row.push(ColumnData::String(value.map(Cow::Owned)));
}
row
}
pub async fn bulk_insert_text_rows<T, F>(
client: &mut SqlServerClient,
staging_table: &str,
rows: &[T],
column_count: usize,
mut convert_row: F,
) -> Result<u64, String>
where
F: FnMut(usize, &T) -> Result<Vec<Option<String>>, String>,
{
if rows.is_empty() {
return Ok(0);
}
if column_count == 0 {
return Err("SQL Server bulk load requires at least one mapped column".to_string());
}
let mut request = client
.bulk_insert(staging_table)
.await
.map_err(|error| format!("SQL Server bulk load initialization failed: {error}"))?;
for (row_index, source_row) in rows.iter().enumerate() {
let row = convert_row(row_index, source_row)?;
if row.len() != column_count {
return Err(format!(
"SQL Server bulk row {} has {} columns; expected {}",
row_index + 1,
row.len(),
column_count
));
}
request
.send(sqlserver_bulk_token_row(row))
.await
.map_err(|error| format!("SQL Server bulk load send failed: {error}"))?;
}
request
.finalize()
.await
.map(|result| result.total())
.map_err(|error| format!("SQL Server bulk load finalize failed: {error}"))
}
pub async fn execute_query_with_max_rows(
client: &mut SqlServerClient,
sql: &str,
@@ -2336,6 +2390,10 @@ fn contains_transaction_control(sql: &str) -> bool {
}
fn requires_simple_query_batch(sql: &str) -> bool {
if creates_local_temp_table(sql) {
return true;
}
let tokens = first_sql_tokens(sql, 4);
if tokens.len() >= 2 && tokens[0].eq_ignore_ascii_case("SET") && tokens[1].eq_ignore_ascii_case("SHOWPLAN_XML") {
return true;
@@ -2359,6 +2417,22 @@ fn requires_simple_query_batch(sql: &str) -> bool {
false
}
fn creates_local_temp_table(sql: &str) -> bool {
if !sql.as_bytes().contains(&b'#') {
return false;
}
let Ok(statements) = Parser::parse_sql(&MsSqlDialect {}, sql) else {
return false;
};
statements.iter().any(|statement| {
let Statement::CreateTable(table) = statement else {
return false;
};
table.name.0.last().and_then(|part| part.as_ident()).is_some_and(|identifier| identifier.value.starts_with('#'))
})
}
fn first_sql_tokens(sql: &str, limit: usize) -> Vec<String> {
let bytes = sql.as_bytes();
let mut tokens = Vec::new();
@@ -2407,7 +2481,7 @@ mod tests {
build_sqlserver_unsafe_type_query, capture_sqlserver_messages, format_sqlserver_numeric,
is_blocking_sqlserver_unsafe_probe_error, is_sqlserver_spatial_column, is_sqlserver_variant_column,
query_result_with_server_messages, requires_simple_query_batch, restore_sqlserver_legacy_probe_output_names,
sqlserver_batch_can_use_execute, sqlserver_cell_to_json, sqlserver_columns_sql,
sqlserver_batch_can_use_execute, sqlserver_bulk_token_row, sqlserver_cell_to_json, sqlserver_columns_sql,
sqlserver_completion_assistant_sql, sqlserver_dml_output_returns_rows, sqlserver_filter_definition_error,
sqlserver_hidden_schema_names, sqlserver_indexes_sql, sqlserver_legacy_indexes_sql, sqlserver_legacy_probe,
sqlserver_legacy_probe_with_nonce, sqlserver_list_objects_sql, sqlserver_list_schemas_sql,
@@ -2423,6 +2497,15 @@ mod tests {
use std::{borrow::Cow, time::Instant};
use tiberius::{Column, ColumnData, ColumnType, IntoSql};
#[test]
fn sqlserver_bulk_token_row_owns_text_and_preserves_nulls() {
let row = sqlserver_bulk_token_row(vec![Some("Tieng Viet".to_string()), None]);
let values = row.iter().collect::<Vec<_>>();
assert!(matches!(&values[0], ColumnData::String(Some(value)) if value.as_ref() == "Tieng Viet"));
assert!(matches!(&values[1], ColumnData::String(None)));
}
#[tokio::test]
async fn sqlserver_ignores_non_info_tiberius_events() {
let (_, messages) = capture_sqlserver_messages(async {
@@ -2588,6 +2671,16 @@ mod tests {
assert!(!requires_simple_query_batch("UPDATE dbo.t SET id = 1;"));
}
#[test]
fn sqlserver_local_temp_table_creation_keeps_session_scoped_query_path() {
assert!(requires_simple_query_batch("CREATE TABLE #stage (id INT);"));
assert!(requires_simple_query_batch("CREATE TABLE [#stage] ([id] INT);"));
assert!(requires_simple_query_batch(
"DECLARE @id INT = 1; CREATE TABLE #stage (id INT); INSERT INTO #stage VALUES (@id);"
));
assert!(!requires_simple_query_batch("CREATE TABLE dbo.stage (id INT);"));
}
#[test]
fn sqlserver_cud_batches_use_execute_for_affected_rows() {
assert!(sqlserver_batch_can_use_execute("UPDATE dbo.users SET active = 0 WHERE id = 1;"));
File diff suppressed because it is too large Load Diff
+297 -63
View File
@@ -23,6 +23,33 @@ const MAX_ORACLE_INSERT_ALL_ROWS: usize = 500;
const MAX_ORACLE_MERGE_ROWS: usize = 500;
const TRANSFER_TARGET_TABLE_LOOKUP_LIMIT: usize = 1000;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct SqlBatchLimits {
max_rows: usize,
target_sql_bytes: usize,
hard_sql_bytes: Option<usize>,
}
impl SqlBatchLimits {
pub(crate) fn for_database(db_type: &DatabaseType, requested_max_rows: usize) -> Self {
let max_rows = requested_max_rows.max(1).min(match db_type {
DatabaseType::SqlServer => MAX_SQLSERVER_INSERT_ROWS,
DatabaseType::Oracle => MAX_ORACLE_INSERT_ALL_ROWS,
_ => usize::MAX,
});
let target_sql_bytes = match db_type {
DatabaseType::CloudflareD1 => crate::db::cloudflare_d1::MAX_SQL_STATEMENT_BYTES,
_ => MAX_TRANSFER_WRITE_SQL_BYTES,
};
Self { max_rows, target_sql_bytes, hard_sql_bytes: None }
}
pub(crate) fn with_hard_sql_bytes(mut self, hard_sql_bytes: Option<usize>) -> Self {
self.hard_sql_bytes = hard_sql_bytes;
self
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq)]
#[serde(rename_all = "camelCase")]
pub enum TransferMode {
@@ -308,12 +335,28 @@ fn is_postgres_integer_like_type(data_type: &str) -> bool {
postgres_integer_bounds(data_type).is_some()
}
pub(crate) fn normalize_postgres_integer_literal(
fn sqlserver_integer_bounds(data_type: &str) -> Option<(i128, i128)> {
let normalized = data_type.trim().to_ascii_lowercase();
match normalized.split(['(', ' ']).next().unwrap_or("") {
"bit" => Some((i128::MIN, i128::MAX)),
"tinyint" => Some((0, i128::from(u8::MAX))),
"smallint" => Some((i128::from(i16::MIN), i128::from(i16::MAX))),
"int" | "integer" => Some((i128::from(i32::MIN), i128::from(i32::MAX))),
"bigint" => Some((i128::from(i64::MIN), i128::from(i64::MAX))),
_ => None,
}
}
pub(crate) fn normalize_integer_literal(
value: &str,
db_type: &DatabaseType,
column_type: Option<&str>,
) -> Option<String> {
let bounds = column_type.filter(|_| is_postgres_transfer_dialect(db_type)).and_then(postgres_integer_bounds)?;
let bounds = match db_type {
db_type if is_postgres_transfer_dialect(db_type) => column_type.and_then(postgres_integer_bounds),
DatabaseType::SqlServer => column_type.and_then(sqlserver_integer_bounds),
_ => None,
}?;
// Excel numeric cells arrive as f64; normalize only an explicit zero fraction so real decimals,
// scientific notation, and values outside the target integer range stay untouched.
@@ -1115,7 +1158,7 @@ pub fn escape_value_typed(val: &serde_json::Value, db_type: &DatabaseType, colum
}
},
serde_json::Value::Number(n) => {
if let Some(integer_literal) = normalize_postgres_integer_literal(&n.to_string(), db_type, column_type) {
if let Some(integer_literal) = normalize_integer_literal(&n.to_string(), db_type, column_type) {
return integer_literal;
}
match db_type {
@@ -1130,7 +1173,7 @@ pub fn escape_value_typed(val: &serde_json::Value, db_type: &DatabaseType, colum
}
}
serde_json::Value::String(s) => {
if let Some(integer_literal) = normalize_postgres_integer_literal(s, db_type, column_type) {
if let Some(integer_literal) = normalize_integer_literal(s, db_type, column_type) {
return integer_literal;
}
if let Some(binary_literal) = format_postgres_binary_sql_literal(s, db_type, column_type) {
@@ -1139,6 +1182,9 @@ pub fn escape_value_typed(val: &serde_json::Value, db_type: &DatabaseType, colum
if let Some(binary_literal) = format_mysql_binary_sql_literal(s, db_type, column_type) {
return binary_literal;
}
if let Some(binary_literal) = format_sqlserver_binary_sql_literal(s, db_type, column_type) {
return binary_literal;
}
if let Some(numeric_literal) = format_mysql_numeric_string_literal(s, db_type, column_type) {
return numeric_literal;
}
@@ -1250,6 +1296,29 @@ fn format_mysql_binary_sql_literal(value: &str, db_type: &DatabaseType, column_t
}
}
fn format_sqlserver_binary_sql_literal(
value: &str,
db_type: &DatabaseType,
column_type: Option<&str>,
) -> Option<String> {
if !matches!(db_type, DatabaseType::SqlServer) {
return None;
}
let column_type = column_type.filter(|column_type| is_binary_transfer_column_type(column_type))?;
let trimmed = value.trim();
if let Some(hex) = trimmed.strip_prefix("0x").or_else(|| trimmed.strip_prefix("0X")) {
if hex.len() % 2 == 0 && hex.as_bytes().iter().all(|byte| byte.is_ascii_hexdigit()) {
return Some(format!("0x{hex}"));
}
}
// SQL Server does not implicitly convert NVARCHAR literals to binary targets.
// Use the target type so text fallback preserves the same Unicode byte encoding
// that a direct typed conversion would produce.
let escaped = value.replace('\'', "''");
Some(format!("CONVERT({column_type}, N'{escaped}')"))
}
fn format_oracle_temporal_sql_literal(
value: &str,
db_type: &DatabaseType,
@@ -1914,21 +1983,71 @@ pub fn generate_insert_typed(
return String::new();
}
let full_table = qualified_table(table, schema, db_type);
let col_list = columns.iter().map(|c| quote_identifier(c, db_type)).collect::<Vec<_>>().join(", ");
let template = InsertSqlTemplate::new(columns, table, schema, db_type);
let value_rows = value_rows_sql(rows, column_types, db_type);
if matches!(db_type, DatabaseType::Oracle) && rows.len() > 1 {
// Oracle 11g does not accept comma-separated multi-row VALUES lists.
let into_rows = value_rows
.iter()
.map(|values| format!("INTO {full_table} ({col_list}) VALUES {values}"))
.collect::<Vec<_>>()
.join("\n");
return format!("INSERT ALL\n{into_rows}\nSELECT 1 FROM dual");
template.build(&value_rows)
}
#[derive(Debug)]
struct InsertSqlTemplate {
standard_prefix: String,
oracle_into_prefix: Option<String>,
}
impl InsertSqlTemplate {
fn new(columns: &[String], table: &str, schema: &str, db_type: &DatabaseType) -> Self {
let full_table = qualified_table(table, schema, db_type);
let col_list = columns.iter().map(|column| quote_identifier(column, db_type)).collect::<Vec<_>>().join(", ");
Self {
standard_prefix: format!("INSERT INTO {full_table} ({col_list}) VALUES\n"),
oracle_into_prefix: matches!(db_type, DatabaseType::Oracle)
.then(|| format!("INTO {full_table} ({col_list}) VALUES ")),
}
}
format!("INSERT INTO {full_table} ({col_list}) VALUES\n{}", value_rows.join(",\n"))
fn build(&self, value_rows: &[String]) -> String {
if value_rows.is_empty() {
return String::new();
}
if let Some(into_prefix) = self.oracle_into_prefix.as_deref().filter(|_| value_rows.len() > 1) {
let mut sql = String::from("INSERT ALL\n");
for (index, values) in value_rows.iter().enumerate() {
if index > 0 {
sql.push('\n');
}
sql.push_str(into_prefix);
sql.push_str(values);
}
sql.push_str("\nSELECT 1 FROM dual");
return sql;
}
let mut sql = self.standard_prefix.clone();
sql.push_str(&value_rows.join(",\n"));
sql
}
fn statement_bytes(&self, value_rows_bytes: usize, row_count: usize, db_type: &DatabaseType) -> usize {
if let Some(into_prefix) = self.oracle_into_prefix.as_deref().filter(|_| row_count > 1) {
return sql_text_bytes("INSERT ALL\n", db_type)
.saturating_add(sql_text_bytes(into_prefix, db_type).saturating_mul(row_count))
.saturating_add(value_rows_bytes)
.saturating_add(sql_text_bytes("\n", db_type).saturating_mul(row_count - 1))
.saturating_add(sql_text_bytes("\nSELECT 1 FROM dual", db_type));
}
sql_text_bytes(&self.standard_prefix, db_type)
.saturating_add(value_rows_bytes)
.saturating_add(sql_text_bytes(",\n", db_type).saturating_mul(row_count.saturating_sub(1)))
}
}
fn sql_text_bytes(sql: &str, db_type: &DatabaseType) -> usize {
if matches!(db_type, DatabaseType::SqlServer) {
sql.encode_utf16().count().saturating_mul(2)
} else {
sql.len()
}
}
fn value_rows_sql(
@@ -2167,53 +2286,55 @@ pub(crate) fn generate_insert_typed_sql_batches(
table: &str,
schema: &str,
db_type: &DatabaseType,
requested_max_rows: usize,
) -> Vec<(String, usize)> {
limits: SqlBatchLimits,
) -> Result<Vec<(String, usize)>, String> {
if rows.is_empty() {
return Vec::new();
return Ok(Vec::new());
}
let max_rows = requested_max_rows.max(1);
let max_sql_bytes = match db_type {
DatabaseType::CloudflareD1 => crate::db::cloudflare_d1::MAX_SQL_STATEMENT_BYTES,
_ => MAX_TRANSFER_WRITE_SQL_BYTES,
};
let max_rows = limits.max_rows.max(1).min(match db_type {
DatabaseType::SqlServer => MAX_SQLSERVER_INSERT_ROWS,
DatabaseType::Oracle => MAX_ORACLE_INSERT_ALL_ROWS,
_ => usize::MAX,
});
let target_sql_bytes = limits.target_sql_bytes.max(1);
let batch_sql_bytes = limits.hard_sql_bytes.map_or(target_sql_bytes, |hard| target_sql_bytes.min(hard));
let template = InsertSqlTemplate::new(columns, table, schema, db_type);
let value_rows = value_rows_sql(rows, column_types, db_type);
let value_row_bytes = value_rows.iter().map(|row| sql_text_bytes(row, db_type)).collect::<Vec<_>>();
let mut statements = Vec::new();
let mut start = 0;
let mut start = 0usize;
// First honor the row limit, then use binary search to find the largest statement that
// also fits the backend byte limit. This avoids generating every intermediate size.
while start < rows.len() {
let max_end = start.saturating_add(max_rows).min(rows.len());
let mut end = max_end;
let mut accepted = generate_insert_typed(columns, column_types, &rows[start..max_end], table, schema, db_type);
if accepted.len() > max_sql_bytes && max_end > start + 1 {
end = start + 1;
accepted = generate_insert_typed(columns, column_types, &rows[start..end], table, schema, db_type);
let mut low = start + 2;
let mut high = max_end;
while low <= high {
let candidate_end = low + (high - low) / 2;
let candidate =
generate_insert_typed(columns, column_types, &rows[start..candidate_end], table, schema, db_type);
if candidate.len() <= max_sql_bytes {
accepted = candidate;
end = candidate_end;
low = candidate_end + 1;
} else {
high = candidate_end - 1;
while start < value_rows.len() {
let mut end = start;
let mut rows_bytes = 0usize;
while end < value_rows.len() && end - start < max_rows {
let single_row_bytes = template.statement_bytes(value_row_bytes[end], 1, db_type);
if let Some(hard_sql_bytes) = limits.hard_sql_bytes {
if single_row_bytes > hard_sql_bytes {
return Err(format!(
"SQL batch row {} requires {} bytes and exceeds the {} byte hard limit",
end + 1,
single_row_bytes,
hard_sql_bytes
));
}
}
let candidate_rows_bytes = rows_bytes.saturating_add(value_row_bytes[end]);
let candidate_row_count = end - start + 1;
let candidate_bytes = template.statement_bytes(candidate_rows_bytes, candidate_row_count, db_type);
if candidate_row_count > 1 && candidate_bytes > batch_sql_bytes {
break;
}
rows_bytes = candidate_rows_bytes;
end += 1;
}
if !accepted.is_empty() {
statements.push((accepted, end - start));
}
statements.push((template.build(&value_rows[start..end]), end - start));
start = end;
}
statements
Ok(statements)
}
#[allow(clippy::too_many_arguments)]
@@ -2226,24 +2347,24 @@ fn generate_transfer_write_sql_batches(
schema: &str,
db_type: &DatabaseType,
pk_columns: &[String],
) -> Vec<String> {
) -> Result<Vec<String>, String> {
if rows.is_empty() {
return Vec::new();
return Ok(Vec::new());
}
if matches!(mode, TransferMode::Append | TransferMode::Overwrite) {
return generate_insert_typed_sql_batches(
return Ok(generate_insert_typed_sql_batches(
columns,
column_types,
rows,
table,
schema,
db_type,
max_transfer_write_rows(db_type, mode),
)
SqlBatchLimits::for_database(db_type, max_transfer_write_rows(db_type, mode)),
)?
.into_iter()
.map(|(sql, _)| sql)
.collect();
.collect());
}
let max_rows = max_transfer_write_rows(db_type, mode);
@@ -2291,7 +2412,7 @@ fn generate_transfer_write_sql_batches(
start = end;
}
statements
Ok(statements)
}
pub fn pagination_sql(
@@ -4107,7 +4228,7 @@ where
&request.target_schema,
target_db_type,
&[],
);
)?;
for (statement_index, batch_sql) in write_statements.iter().enumerate() {
execute_on_pool(state, target_pool_key, batch_sql).await.map_err(|e| {
format!(
@@ -4452,7 +4573,7 @@ where
&request.target_schema,
target_db_type,
&pk_columns,
);
)?;
for (statement_index, batch_sql) in write_statements.iter().enumerate() {
execute_transfer_write_statement(
state,
@@ -6313,6 +6434,34 @@ mod tests {
assert_eq!(sql, "INSERT INTO [dbo].[flags] ([enabled], [deleted]) VALUES\n(1, 0)");
}
#[test]
fn sqlserver_insert_formats_prefixed_hex_for_varbinary_columns() {
let sql = generate_insert_typed(
&[String::from("payload"), String::from("note")],
&[Some(String::from("varbinary(max)")), Some(String::from("nvarchar(64)"))],
&[vec![json!("0x0001ABff"), json!("0x0001ABff")]],
"files",
"dbo",
&DatabaseType::SqlServer,
);
assert_eq!(sql, "INSERT INTO [dbo].[files] ([payload], [note]) VALUES\n(0x0001ABff, N'0x0001ABff')");
}
#[test]
fn sqlserver_insert_explicitly_converts_plain_text_for_varbinary_columns() {
let sql = generate_insert_typed(
&[String::from("payload")],
&[Some(String::from("varbinary(max)"))],
&[vec![json!("O'Brien")]],
"files",
"dbo",
&DatabaseType::SqlServer,
);
assert_eq!(sql, "INSERT INTO [dbo].[files] ([payload]) VALUES\n(CONVERT(varbinary(max), N'O''Brien'))");
}
#[test]
fn dameng_insert_formats_bit_booleans_as_numeric_literals() {
let sql = generate_insert_typed(
@@ -6586,7 +6735,8 @@ SELECT 1 FROM dual"#
"APP",
&DatabaseType::Oracle,
&[],
);
)
.unwrap();
assert_eq!(statements.len(), 2);
assert_eq!(statements[0].matches("\nINTO ").count(), MAX_ORACLE_INSERT_ALL_ROWS);
@@ -6606,12 +6756,95 @@ SELECT 1 FROM dual"#
"",
&DatabaseType::Mysql,
&[],
);
)
.unwrap();
assert!(statements.len() > 1);
assert!(statements.iter().all(|sql| sql.starts_with("INSERT INTO `events`")));
}
#[test]
fn mysql_sql_batch_allows_one_row_over_soft_target() {
let rows = vec![vec![json!("x".repeat(256))]];
let limits = SqlBatchLimits { max_rows: 100, target_sql_bytes: 128, hard_sql_bytes: Some(1024) };
let batches = generate_insert_typed_sql_batches(
&[String::from("payload")],
&[Some(String::from("text"))],
&rows,
"events",
"",
&DatabaseType::Mysql,
limits,
)
.unwrap();
assert_eq!(batches.len(), 1);
assert_eq!(batches[0].1, 1);
}
#[test]
fn mysql_sql_batch_rejects_one_row_over_known_hard_limit() {
let rows = vec![vec![json!("x".repeat(256))]];
let limits = SqlBatchLimits { max_rows: 100, target_sql_bytes: 128, hard_sql_bytes: Some(200) };
let error = generate_insert_typed_sql_batches(
&[String::from("payload")],
&[Some(String::from("text"))],
&rows,
"events",
"",
&DatabaseType::Mysql,
limits,
)
.unwrap_err();
assert!(error.contains("row 1"));
assert!(error.contains("200 byte hard limit"));
}
#[test]
fn sqlserver_insert_batches_enforce_values_row_limit() {
let rows = (0..(MAX_SQLSERVER_INSERT_ROWS + 1)).map(|index| vec![json!(index)]).collect::<Vec<_>>();
let batches = generate_insert_typed_sql_batches(
&[String::from("id")],
&[Some(String::from("int"))],
&rows,
"events",
"dbo",
&DatabaseType::SqlServer,
SqlBatchLimits::for_database(&DatabaseType::SqlServer, rows.len()),
)
.unwrap();
assert_eq!(batches.iter().map(|(_, row_count)| *row_count).collect::<Vec<_>>(), vec![1000, 1]);
}
#[test]
fn sqlserver_insert_batches_measure_unicode_sql_as_utf16() {
let rows = (0..2).map(|_| vec![json!("x".repeat(140 * 1024))]).collect::<Vec<_>>();
let batches = generate_insert_typed_sql_batches(
&[String::from("payload")],
&[Some(String::from("nvarchar(max)"))],
&rows,
"events",
"dbo",
&DatabaseType::SqlServer,
SqlBatchLimits::for_database(&DatabaseType::SqlServer, rows.len()),
)
.unwrap();
assert_eq!(batches.iter().map(|(_, row_count)| *row_count).collect::<Vec<_>>(), vec![1, 1]);
}
#[test]
fn sqlserver_sql_byte_count_uses_utf16_code_units() {
assert_eq!(sql_text_bytes("AA\u{8d8a}\u{1f600}", &DatabaseType::SqlServer), 10);
assert_eq!(sql_text_bytes("AA\u{8d8a}\u{1f600}", &DatabaseType::Postgres), 9);
}
#[test]
fn transfer_write_sql_batches_keep_existing_upsert_sql_shape() {
let statements = generate_transfer_write_sql_batches(
@@ -6623,7 +6856,8 @@ SELECT 1 FROM dual"#
"",
&DatabaseType::Mysql,
&[String::from("id")],
);
)
.unwrap();
assert_eq!(statements.len(), 1);
assert!(statements[0].contains("ON DUPLICATE KEY UPDATE"));
@@ -4,10 +4,16 @@ use dbx_core::query_result_export::{export_query_result_core, ExportStatus, Quer
use dbx_core::sql::{SqlFileRequest, SqlFileStatus};
use dbx_core::sql_file_import::execute_sql_file_content;
use dbx_core::storage::Storage;
use dbx_core::table_import::{
build_import_insert_batches, import_table_file_core, parse_delimited_file_with_options, TableImportColumnMapping,
TableImportMode, TableImportParseOptions, TableImportRequest, TableImportSourceFormat, TableImportStatus,
};
use dbx_core::table_structure_sql::{
build_table_structure_change_sql, ColumnInfo, EditableStructureColumn, TableStructureSqlOptions,
};
use std::sync::atomic::{AtomicBool, Ordering};
use dbx_core::xlsx_export::{build_xlsx_workbook, XlsxWorksheetData};
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio_util::sync::CancellationToken;
@@ -68,6 +74,740 @@ fn live_sqlserver_config(id: &str, database: &str) -> dbx_core::models::connecti
}
}
async fn live_sqlserver_import_state(
connection_id: &str,
database: &str,
suffix: &str,
) -> (AppState, String, std::path::PathBuf) {
let dir = std::env::temp_dir().join(format!("dbx-live-sqlserver-import-{suffix}"));
std::fs::create_dir_all(&dir).expect("create live import directory");
let storage = Storage::open(&dir.join("storage.db")).await.expect("open live import storage");
let state = AppState::new(storage);
let config = live_sqlserver_config(connection_id, database);
state.configs.write().await.insert(connection_id.to_string(), config);
let pool_key =
state.get_or_create_pool(connection_id, Some(database)).await.expect("connect live SQL Server import pool");
(state, pool_key, dir)
}
fn live_sqlserver_import_mapping(source: &str, target: &str) -> TableImportColumnMapping {
TableImportColumnMapping {
source_column: source.to_string(),
target_column: target.to_string(),
target_data_type: None,
}
}
fn live_sqlserver_import_request(
connection_id: &str,
database: &str,
table: &str,
file_path: &std::path::Path,
mappings: Vec<TableImportColumnMapping>,
mode: TableImportMode,
) -> TableImportRequest {
TableImportRequest {
import_id: format!("live-sqlserver-import-{}", uuid::Uuid::new_v4().simple()),
connection_id: connection_id.to_string(),
database: database.to_string(),
schema: "dbo".to_string(),
table: table.to_string(),
file_path: file_path.to_string_lossy().to_string(),
source_ref: None,
source_format: Some(TableImportSourceFormat::Csv),
parse_options: TableImportParseOptions::default(),
mappings,
mode,
create_table: false,
batch_size: 500,
date_time_format: None,
prepared_source: None,
retain_source: false,
}
}
async fn run_live_sqlserver_import(
state: &AppState,
pool_key: &str,
request: &TableImportRequest,
) -> Result<dbx_core::table_import::TableImportSummary, String> {
import_table_file_core(state, request, &DatabaseType::SqlServer, pool_key, |_| Box::pin(async { false }), |_| {})
.await
}
fn live_sqlserver_matrix_column_types() -> Vec<(String, String)> {
[
("id", "int"),
("code", "nvarchar(40)"),
("nullable_text", "nvarchar(100)"),
("amount", "decimal(38,10)"),
("occurred_at", "datetime2(7)"),
("offset_at", "datetimeoffset(7)"),
("event_id", "uniqueidentifier"),
("document", "xml"),
("payload", "varbinary(max)"),
]
.into_iter()
.map(|(name, data_type)| (name.to_string(), data_type.to_string()))
.collect()
}
fn live_sqlserver_matrix_mappings(include_identity: bool) -> Vec<TableImportColumnMapping> {
let columns =
["id", "code", "nullable_text", "amount", "occurred_at", "offset_at", "event_id", "document", "payload"];
columns
.into_iter()
.filter(|column| include_identity || *column != "id")
.map(|column| live_sqlserver_import_mapping(column, column))
.collect()
}
async fn run_live_sqlserver_generated_insert(
client: &mut dbx_core::db::sqlserver::SqlServerClient,
table: &str,
file_path: &std::path::Path,
parse_options: &TableImportParseOptions,
include_identity: bool,
) -> Result<usize, String> {
let parsed = parse_delimited_file_with_options(
&file_path.to_string_lossy(),
TableImportSourceFormat::Csv,
parse_options,
usize::MAX,
)?;
let batches = build_import_insert_batches(
&parsed,
&live_sqlserver_matrix_mappings(include_identity),
&live_sqlserver_matrix_column_types(),
table,
"dbo",
&DatabaseType::SqlServer,
500,
)?;
if include_identity {
let rows_imported = batches.iter().map(|batch| batch.row_count).sum();
let statements = batches.into_iter().map(|batch| batch.sql).collect::<Vec<_>>().join(";\n");
let sql =
format!("SET IDENTITY_INSERT [dbo].[{table}] ON;\n{statements};\nSET IDENTITY_INSERT [dbo].[{table}] OFF");
dbx_core::db::sqlserver::execute_batch(client, &sql).await?;
return Ok(rows_imported);
}
let mut rows_imported = 0;
for batch in batches {
dbx_core::db::sqlserver::execute_batch(client, &batch.sql).await?;
rows_imported += batch.row_count;
}
Ok(rows_imported)
}
fn live_sqlserver_matrix_table_ddl(table: &str, audit_table: &str) -> String {
format!(
"CREATE TABLE [dbo].[{table}] (\
[id] INT IDENTITY(1,1) NOT NULL PRIMARY KEY, \
[code] NVARCHAR(40) NOT NULL UNIQUE, \
[nullable_text] NVARCHAR(100) NULL, \
[amount] DECIMAL(38,10) NOT NULL CHECK ([amount] > 0), \
[occurred_at] DATETIME2(7) NOT NULL, \
[offset_at] DATETIMEOFFSET(7) NOT NULL, \
[event_id] UNIQUEIDENTIFIER NOT NULL, \
[document] XML NULL, \
[payload] VARBINARY(MAX) NULL, \
[default_text] NVARCHAR(40) NOT NULL DEFAULT N'defaulted'); \
CREATE TABLE [dbo].[{audit_table}] (\
[target_id] INT NOT NULL, [code] NVARCHAR(40) NOT NULL, [default_text] NVARCHAR(40) NOT NULL);"
)
}
fn live_sqlserver_matrix_trigger_ddl(table: &str, audit_table: &str, trigger: &str) -> String {
format!(
"CREATE TRIGGER [dbo].[{trigger}] ON [dbo].[{table}] AFTER INSERT AS \
INSERT INTO [dbo].[{audit_table}] ([target_id], [code], [default_text]) \
SELECT [id], [code], [default_text] FROM inserted;"
)
}
async fn live_sqlserver_matrix_rows(
client: &mut dbx_core::db::sqlserver::SqlServerClient,
table: &str,
) -> dbx_core::db::QueryResult {
dbx_core::db::sqlserver::execute_query(
client,
&format!(
"SELECT CONVERT(VARCHAR(12), [id]), [code], \
CASE WHEN [nullable_text] IS NULL THEN N'<NULL>' ELSE N'<' + [nullable_text] + N'>' END, \
CONVERT(VARCHAR(50), [amount]), CONVERT(VARCHAR(33), [occurred_at], 126), \
CONVERT(VARCHAR(48), [offset_at], 127), CONVERT(VARCHAR(36), [event_id]), \
CASE WHEN [document] IS NULL THEN N'<NULL>' ELSE CONVERT(NVARCHAR(MAX), [document]) END, \
CASE WHEN [payload] IS NULL THEN N'<NULL>' ELSE sys.fn_varbintohexstr([payload]) END, \
[default_text] FROM [dbo].[{table}] ORDER BY [id]"
),
)
.await
.expect("query SQL Server import matrix rows")
}
async fn live_sqlserver_table_count(client: &mut dbx_core::db::sqlserver::SqlServerClient, table: &str) -> i64 {
let result = dbx_core::db::sqlserver::execute_query(client, &format!("SELECT COUNT_BIG(*) FROM [dbo].[{table}]"))
.await
.expect("count SQL Server matrix rows");
result.rows[0][0].as_i64().expect("SQL Server COUNT_BIG result")
}
#[tokio::test]
#[ignore = "requires DBX_LIVE_SQLSERVER_HOST/PORT/USER/PASSWORD pointing at a writable SQL Server database"]
async fn live_sqlserver_bulk_matches_generated_insert_type_and_constraint_matrix() {
let database = std::env::var("DBX_LIVE_SQLSERVER_DATABASE").unwrap_or_else(|_| "tempdb".to_string());
let suffix = uuid::Uuid::new_v4().simple().to_string();
let connection_id = format!("live-sqlserver-matrix-{suffix}");
let bulk_table = format!("dbx_bulk_matrix_{suffix}");
let generated_table = format!("dbx_insert_matrix_{suffix}");
let bulk_audit = format!("dbx_bulk_matrix_audit_{suffix}");
let generated_audit = format!("dbx_insert_matrix_audit_{suffix}");
let bulk_trigger = format!("dbx_bulk_matrix_trigger_{suffix}");
let generated_trigger = format!("dbx_insert_matrix_trigger_{suffix}");
let mut client = dbx_core::db::sqlserver::connect(
&std::env::var("DBX_LIVE_SQLSERVER_HOST").unwrap_or_else(|_| "127.0.0.1".to_string()),
std::env::var("DBX_LIVE_SQLSERVER_PORT").ok().and_then(|value| value.parse().ok()).unwrap_or(1433),
&std::env::var("DBX_LIVE_SQLSERVER_USER").unwrap_or_else(|_| "sa".to_string()),
&std::env::var("DBX_LIVE_SQLSERVER_PASSWORD").expect("DBX_LIVE_SQLSERVER_PASSWORD"),
Some(&database),
None,
Duration::from_secs(10),
)
.await
.expect("connect SQL Server");
for ddl in [
live_sqlserver_matrix_table_ddl(&bulk_table, &bulk_audit),
live_sqlserver_matrix_table_ddl(&generated_table, &generated_audit),
live_sqlserver_matrix_trigger_ddl(&bulk_table, &bulk_audit, &bulk_trigger),
live_sqlserver_matrix_trigger_ddl(&generated_table, &generated_audit, &generated_trigger),
] {
dbx_core::db::sqlserver::execute_batch(&mut client, &ddl)
.await
.expect("create SQL Server import matrix objects");
}
let (state, pool_key, dir) = live_sqlserver_import_state(&connection_id, &database, &suffix).await;
let default_options = TableImportParseOptions::default();
let empty_string_options =
TableImportParseOptions { empty_string_as_null: Some(false), ..TableImportParseOptions::default() };
let null_csv = dir.join("matrix-null.csv");
std::fs::write(
&null_csv,
"code,nullable_text,amount,occurred_at,offset_at,event_id,document,payload\n\
auto-null,,12345678901234567890.1234567890,2026-07-27T12:34:56.1234567,2026-07-27T12:34:56.1234567+08:00,11111111-2222-3333-4444-555555555555,,\n",
)
.expect("write SQL Server NULL matrix CSV");
let null_request = live_sqlserver_import_request(
&connection_id,
&database,
&bulk_table,
&null_csv,
live_sqlserver_matrix_mappings(false),
TableImportMode::Append,
);
let null_bulk = run_live_sqlserver_import(&state, &pool_key, &null_request)
.await
.expect("bulk import SQL Server NULL matrix row");
let null_generated =
run_live_sqlserver_generated_insert(&mut client, &generated_table, &null_csv, &default_options, false)
.await
.expect("generated INSERT SQL Server NULL matrix row");
let empty_csv = dir.join("matrix-empty-string.csv");
std::fs::write(
&empty_csv,
"code,nullable_text,amount,occurred_at,offset_at,event_id,document,payload\n\
empty,,0.0000000001,2026-07-28T01:02:03.0000001,2026-07-28T01:02:03.0000001-05:30,aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee,<root><value>empty-text</value></root>,0x00FF10\n",
)
.expect("write SQL Server empty-string matrix CSV");
let mut empty_request = live_sqlserver_import_request(
&connection_id,
&database,
&bulk_table,
&empty_csv,
live_sqlserver_matrix_mappings(false),
TableImportMode::Append,
);
empty_request.parse_options = empty_string_options.clone();
let empty_bulk = run_live_sqlserver_import(&state, &pool_key, &empty_request)
.await
.expect("bulk import SQL Server empty-string matrix row");
let empty_generated =
run_live_sqlserver_generated_insert(&mut client, &generated_table, &empty_csv, &empty_string_options, false)
.await
.expect("generated INSERT SQL Server empty-string matrix row");
let identity_csv = dir.join("matrix-explicit-identity.csv");
std::fs::write(
&identity_csv,
"id,code,nullable_text,amount,occurred_at,offset_at,event_id,document,payload\n\
42,explicit,identity,9999999999999999999999999999.9999999999,2026-07-29T23:59:59.9999999,2026-07-29T23:59:59.9999999+13:45,01234567-89ab-cdef-0123-456789abcdef,<root><value>identity</value></root>,0xABCDEF\n",
)
.expect("write SQL Server explicit-identity matrix CSV");
let identity_request = live_sqlserver_import_request(
&connection_id,
&database,
&bulk_table,
&identity_csv,
live_sqlserver_matrix_mappings(true),
TableImportMode::Append,
);
let identity_bulk = run_live_sqlserver_import(&state, &pool_key, &identity_request)
.await
.expect("bulk import SQL Server explicit-identity matrix row");
let identity_generated =
run_live_sqlserver_generated_insert(&mut client, &generated_table, &identity_csv, &default_options, true)
.await
.expect("generated INSERT SQL Server explicit-identity matrix row");
let bulk_rows = live_sqlserver_matrix_rows(&mut client, &bulk_table).await;
let generated_rows = live_sqlserver_matrix_rows(&mut client, &generated_table).await;
let bulk_audit_rows = dbx_core::db::sqlserver::execute_query(
&mut client,
&format!("SELECT [target_id], [code], [default_text] FROM [dbo].[{bulk_audit}] ORDER BY [target_id]"),
)
.await
.expect("query SQL Server bulk audit rows");
let generated_audit_rows = dbx_core::db::sqlserver::execute_query(
&mut client,
&format!("SELECT [target_id], [code], [default_text] FROM [dbo].[{generated_audit}] ORDER BY [target_id]"),
)
.await
.expect("query SQL Server generated INSERT audit rows");
let invalid_cases = [
(
"unique",
"code,nullable_text,amount,occurred_at,offset_at,event_id,document,payload\n\
unique-first,value,1.0000000000,2026-08-01T00:00:00,2026-08-01T00:00:00+00:00,10000000-0000-0000-0000-000000000001,<ok/>,0x01\n\
auto-null,value,2.0000000000,2026-08-01T00:00:01,2026-08-01T00:00:01+00:00,10000000-0000-0000-0000-000000000002,<ok/>,0x02\n",
),
(
"check",
"code,nullable_text,amount,occurred_at,offset_at,event_id,document,payload\n\
check-first,value,1.0000000000,2026-08-02T00:00:00,2026-08-02T00:00:00+00:00,20000000-0000-0000-0000-000000000001,<ok/>,0x01\n\
check-invalid,value,-0.0000000001,2026-08-02T00:00:01,2026-08-02T00:00:01+00:00,20000000-0000-0000-0000-000000000002,<ok/>,0x02\n",
),
(
"not-null",
"code,nullable_text,amount,occurred_at,offset_at,event_id,document,payload\n\
not-null-first,value,1.0000000000,2026-08-03T00:00:00,2026-08-03T00:00:00+00:00,30000000-0000-0000-0000-000000000001,<ok/>,0x01\n\
,value,2.0000000000,2026-08-03T00:00:01,2026-08-03T00:00:01+00:00,30000000-0000-0000-0000-000000000002,<ok/>,0x02\n",
),
(
"decimal-overflow",
"code,nullable_text,amount,occurred_at,offset_at,event_id,document,payload\n\
decimal-first,value,1.0000000000,2026-08-04T00:00:00,2026-08-04T00:00:00+00:00,40000000-0000-0000-0000-000000000001,<ok/>,0x01\n\
decimal-invalid,value,10000000000000000000000000000.0000000000,2026-08-04T00:00:01,2026-08-04T00:00:01+00:00,40000000-0000-0000-0000-000000000002,<ok/>,0x02\n",
),
(
"uuid",
"code,nullable_text,amount,occurred_at,offset_at,event_id,document,payload\n\
uuid-first,value,1.0000000000,2026-08-05T00:00:00,2026-08-05T00:00:00+00:00,50000000-0000-0000-0000-000000000001,<ok/>,0x01\n\
uuid-invalid,value,2.0000000000,2026-08-05T00:00:01,2026-08-05T00:00:01+00:00,not-a-uuid,<ok/>,0x02\n",
),
(
"xml",
"code,nullable_text,amount,occurred_at,offset_at,event_id,document,payload\n\
xml-first,value,1.0000000000,2026-08-06T00:00:00,2026-08-06T00:00:00+00:00,60000000-0000-0000-0000-000000000001,<ok/>,0x01\n\
xml-invalid,value,2.0000000000,2026-08-06T00:00:01,2026-08-06T00:00:01+00:00,60000000-0000-0000-0000-000000000002,<unclosed>,0x02\n",
),
];
let stable_bulk_count = live_sqlserver_table_count(&mut client, &bulk_table).await;
let stable_generated_count = live_sqlserver_table_count(&mut client, &generated_table).await;
let mut rejection_results = Vec::new();
for (case, csv) in invalid_cases {
let path = dir.join(format!("matrix-invalid-{case}.csv"));
std::fs::write(&path, csv).expect("write SQL Server invalid matrix CSV");
let request = live_sqlserver_import_request(
&connection_id,
&database,
&bulk_table,
&path,
live_sqlserver_matrix_mappings(false),
TableImportMode::Append,
);
let bulk_result = run_live_sqlserver_import(&state, &pool_key, &request).await;
let generated_result =
run_live_sqlserver_generated_insert(&mut client, &generated_table, &path, &default_options, false).await;
let bulk_count = live_sqlserver_table_count(&mut client, &bulk_table).await;
let generated_count = live_sqlserver_table_count(&mut client, &generated_table).await;
rejection_results.push((case, bulk_result, generated_result, bulk_count, generated_count));
}
let final_bulk_audit_count = live_sqlserver_table_count(&mut client, &bulk_audit).await;
let final_generated_audit_count = live_sqlserver_table_count(&mut client, &generated_audit).await;
let cleanup = format!(
"DROP TRIGGER IF EXISTS [dbo].[{bulk_trigger}]; \
DROP TRIGGER IF EXISTS [dbo].[{generated_trigger}]; \
DROP TABLE IF EXISTS [dbo].[{bulk_audit}]; DROP TABLE IF EXISTS [dbo].[{generated_audit}]; \
DROP TABLE IF EXISTS [dbo].[{bulk_table}]; DROP TABLE IF EXISTS [dbo].[{generated_table}];"
);
let _ = dbx_core::db::sqlserver::execute_batch(&mut client, &cleanup).await;
state.remove_connection_pools_detached(&connection_id).await;
let _ = std::fs::remove_dir_all(&dir);
assert_eq!(null_bulk.rows_imported, null_generated);
assert_eq!(empty_bulk.rows_imported, empty_generated);
assert_eq!(identity_bulk.rows_imported, identity_generated);
assert_eq!(bulk_rows.rows, generated_rows.rows, "Bulk and generated INSERT values differ");
assert_eq!(bulk_rows.rows.len(), 3);
assert_eq!(bulk_rows.rows[0][2], serde_json::json!("<NULL>"));
assert_eq!(bulk_rows.rows[0][3], serde_json::json!("12345678901234567890.1234567890"));
assert_eq!(bulk_rows.rows[0][5], serde_json::json!("2026-07-27T04:34:56.1234567Z"));
assert_eq!(bulk_rows.rows[0][7], serde_json::json!("<NULL>"));
assert_eq!(bulk_rows.rows[0][8], serde_json::json!("<NULL>"));
assert_eq!(bulk_rows.rows[1][2], serde_json::json!("<>"));
assert_eq!(bulk_rows.rows[1][3], serde_json::json!("0.0000000001"));
assert_eq!(bulk_rows.rows[1][5], serde_json::json!("2026-07-28T06:32:03.0000001Z"));
assert_eq!(bulk_rows.rows[1][8], serde_json::json!("0x00ff10"));
assert_eq!(bulk_rows.rows[2][0], serde_json::json!("42"));
assert_eq!(bulk_rows.rows[2][3], serde_json::json!("9999999999999999999999999999.9999999999"));
assert_eq!(bulk_rows.rows[2][5], serde_json::json!("2026-07-29T10:14:59.9999999Z"));
assert!(bulk_rows.rows.iter().all(|row| row[9] == serde_json::json!("defaulted")));
assert_eq!(bulk_audit_rows.rows, generated_audit_rows.rows, "trigger side effects differ");
assert_eq!(bulk_audit_rows.rows.len(), 3);
assert_eq!(final_bulk_audit_count, 3, "Bulk constraint failures left trigger side effects");
assert_eq!(final_generated_audit_count, 3, "generated INSERT failures left trigger side effects");
for (case, bulk_result, generated_result, bulk_count, generated_count) in rejection_results {
assert!(bulk_result.is_err(), "Bulk path unexpectedly accepted {case} case");
assert!(generated_result.is_err(), "generated INSERT unexpectedly accepted {case} case");
assert_eq!(bulk_count, stable_bulk_count, "Bulk path partially wrote {case} case");
assert_eq!(generated_count, stable_generated_count, "generated INSERT partially wrote {case} case");
}
}
#[tokio::test]
#[ignore = "requires DBX_LIVE_SQLSERVER_HOST/PORT/USER/PASSWORD pointing at a writable SQL Server database"]
async fn live_sqlserver_bulk_imports_zero_fraction_xlsx_numbers_into_bigint() {
let database = std::env::var("DBX_LIVE_SQLSERVER_DATABASE").unwrap_or_else(|_| "tempdb".to_string());
let suffix = uuid::Uuid::new_v4().simple().to_string();
let connection_id = format!("live-sqlserver-xlsx-integer-{suffix}");
let table = format!("dbx_bulk_xlsx_integer_{suffix}");
let mut client = dbx_core::db::sqlserver::connect(
&std::env::var("DBX_LIVE_SQLSERVER_HOST").unwrap_or_else(|_| "127.0.0.1".to_string()),
std::env::var("DBX_LIVE_SQLSERVER_PORT").ok().and_then(|value| value.parse().ok()).unwrap_or(1433),
&std::env::var("DBX_LIVE_SQLSERVER_USER").unwrap_or_else(|_| "sa".to_string()),
&std::env::var("DBX_LIVE_SQLSERVER_PASSWORD").expect("DBX_LIVE_SQLSERVER_PASSWORD"),
Some(&database),
None,
Duration::from_secs(10),
)
.await
.expect("connect SQL Server");
dbx_core::db::sqlserver::execute_batch(
&mut client,
&format!("CREATE TABLE [dbo].[{table}] ([id] BIGINT NOT NULL, [label] NVARCHAR(40) NOT NULL)"),
)
.await
.expect("create SQL Server XLSX integer table");
let (state, pool_key, dir) = live_sqlserver_import_state(&connection_id, &database, &suffix).await;
let xlsx = build_xlsx_workbook(&XlsxWorksheetData {
sheet_name: Some("Numbers".to_string()),
columns: vec!["id".to_string(), "label".to_string()],
column_types: Vec::new(),
rows: vec![vec![serde_json::json!(1.0), serde_json::json!("xlsx")]],
numeric_column_right_align: false,
})
.expect("build SQL Server XLSX integer fixture");
let path = dir.join("zero-fraction-integer.xlsx");
std::fs::write(&path, xlsx).expect("write SQL Server XLSX integer fixture");
let mut request = live_sqlserver_import_request(
&connection_id,
&database,
&table,
&path,
vec![live_sqlserver_import_mapping("id", "id"), live_sqlserver_import_mapping("label", "label")],
TableImportMode::Append,
);
request.source_format = Some(TableImportSourceFormat::Excel);
let result = run_live_sqlserver_import(&state, &pool_key, &request).await;
let rows =
dbx_core::db::sqlserver::execute_query(&mut client, &format!("SELECT [id], [label] FROM [dbo].[{table}]"))
.await
.expect("query SQL Server XLSX integer row");
let _ = dbx_core::db::sqlserver::execute_batch(&mut client, &format!("DROP TABLE IF EXISTS [dbo].[{table}]")).await;
state.remove_connection_pools_detached(&connection_id).await;
let _ = std::fs::remove_dir_all(&dir);
assert_eq!(result.expect("bulk import SQL Server XLSX integer row").rows_imported, 1);
assert_eq!(rows.rows, vec![vec![serde_json::json!(1), serde_json::json!("xlsx")]]);
}
#[tokio::test]
#[ignore = "requires DBX_LIVE_SQLSERVER_HOST/PORT/USER/PASSWORD pointing at a writable SQL Server database"]
async fn live_sqlserver_table_import_bulk_preserves_target_semantics() {
let database = std::env::var("DBX_LIVE_SQLSERVER_DATABASE").unwrap_or_else(|_| "tempdb".to_string());
let suffix = uuid::Uuid::new_v4().simple().to_string();
let connection_id = format!("live-sqlserver-import-{suffix}");
let table = format!("dbx_bulk_target_{suffix}");
let audit_table = format!("dbx_bulk_audit_{suffix}");
let trigger = format!("dbx_bulk_trigger_{suffix}");
let mut setup_client = dbx_core::db::sqlserver::connect(
&std::env::var("DBX_LIVE_SQLSERVER_HOST").unwrap_or_else(|_| "127.0.0.1".to_string()),
std::env::var("DBX_LIVE_SQLSERVER_PORT").ok().and_then(|value| value.parse().ok()).unwrap_or(1433),
&std::env::var("DBX_LIVE_SQLSERVER_USER").unwrap_or_else(|_| "sa".to_string()),
&std::env::var("DBX_LIVE_SQLSERVER_PASSWORD").expect("DBX_LIVE_SQLSERVER_PASSWORD"),
Some(&database),
None,
Duration::from_secs(10),
)
.await
.expect("connect SQL Server");
dbx_core::db::sqlserver::execute_batch(
&mut setup_client,
&format!(
"CREATE TABLE [dbo].[{table}] (\
[id] INT IDENTITY(1,1) NOT NULL PRIMARY KEY, \
[code] NVARCHAR(40) NOT NULL UNIQUE, \
[occurred_at] DATETIME2(7) NOT NULL, \
[amount] DECIMAL(38,10) NOT NULL, \
[name] NVARCHAR(100) NOT NULL, \
[payload] VARBINARY(MAX) NULL, \
[created_at] DATETIME2(7) NOT NULL DEFAULT SYSUTCDATETIME()); \
CREATE TABLE [dbo].[{audit_table}] ([target_id] INT NOT NULL, [name] NVARCHAR(100) NOT NULL);"
),
)
.await
.expect("create bulk target tables");
dbx_core::db::sqlserver::execute_batch(
&mut setup_client,
&format!(
"CREATE TRIGGER [dbo].[{trigger}] ON [dbo].[{table}] AFTER INSERT AS \
INSERT INTO [dbo].[{audit_table}] ([target_id], [name]) SELECT [id], [name] FROM inserted"
),
)
.await
.expect("create bulk target trigger");
let (state, pool_key, dir) = live_sqlserver_import_state(&connection_id, &database, &suffix).await;
let generated_identity_csv = dir.join("generated-identity.csv");
std::fs::write(
&generated_identity_csv,
"code,occurred_at,amount,name,payload\nauto,2026-07-27T12:34:56.1234567,12345678901234567890.1234567890,\u{8d8a}\u{5357}\u{82b1},0x00FF10\n",
)
.expect("write generated identity CSV");
let generated_identity_request = live_sqlserver_import_request(
&connection_id,
&database,
&table,
&generated_identity_csv,
["code", "occurred_at", "amount", "name", "payload"]
.into_iter()
.map(|column| live_sqlserver_import_mapping(column, column))
.collect(),
TableImportMode::Append,
);
let generated_summary = run_live_sqlserver_import(&state, &pool_key, &generated_identity_request)
.await
.expect("bulk import with generated identity");
let explicit_identity_csv = dir.join("explicit-identity.csv");
std::fs::write(
&explicit_identity_csv,
"id,code,occurred_at,amount,name,payload\n42,explicit,2026-07-28T01:02:03.0000000,1.0000000000,Tieng Viet,0xABCDEF\n",
)
.expect("write explicit identity CSV");
let explicit_identity_request = live_sqlserver_import_request(
&connection_id,
&database,
&table,
&explicit_identity_csv,
["id", "code", "occurred_at", "amount", "name", "payload"]
.into_iter()
.map(|column| live_sqlserver_import_mapping(column, column))
.collect(),
TableImportMode::Append,
);
let explicit_summary = run_live_sqlserver_import(&state, &pool_key, &explicit_identity_request)
.await
.expect("bulk import with explicit identity");
let plain_binary_csv = dir.join("plain-binary.csv");
std::fs::write(
&plain_binary_csv,
"code,occurred_at,amount,name,payload\nplain,2026-07-28T02:03:04.0000000,2.0000000000,plain binary,plain\n",
)
.expect("write plain binary CSV");
let plain_binary_request = live_sqlserver_import_request(
&connection_id,
&database,
&table,
&plain_binary_csv,
["code", "occurred_at", "amount", "name", "payload"]
.into_iter()
.map(|column| live_sqlserver_import_mapping(column, column))
.collect(),
TableImportMode::Append,
);
let plain_binary_summary = run_live_sqlserver_import(&state, &pool_key, &plain_binary_request)
.await
.expect("SQL fallback import with plain-text varbinary input");
let duplicate_csv = dir.join("duplicate.csv");
std::fs::write(
&duplicate_csv,
"code,occurred_at,amount,name,payload\nauto,2026-07-29T00:00:00,2.0000000000,duplicate,0x01\n",
)
.expect("write duplicate CSV");
let duplicate_request = live_sqlserver_import_request(
&connection_id,
&database,
&table,
&duplicate_csv,
["code", "occurred_at", "amount", "name", "payload"]
.into_iter()
.map(|column| live_sqlserver_import_mapping(column, column))
.collect(),
TableImportMode::Append,
);
let duplicate_error = run_live_sqlserver_import(&state, &pool_key, &duplicate_request).await;
let rows = dbx_core::db::sqlserver::execute_query(
&mut setup_client,
&format!(
"SELECT CONVERT(VARCHAR(12), [id]), [code], CONVERT(VARCHAR(33), [occurred_at], 126), \
CONVERT(VARCHAR(50), [amount]), [name], sys.fn_varbintohexstr([payload]), \
CASE WHEN [created_at] IS NULL THEN N'missing' ELSE N'set' END, \
CASE WHEN [code] = N'plain' THEN CONVERT(NVARCHAR(100), [payload]) END \
FROM [dbo].[{table}] ORDER BY [id]"
),
)
.await
.expect("verify bulk target rows");
let audit_count = dbx_core::db::sqlserver::execute_query(
&mut setup_client,
&format!("SELECT COUNT(*) FROM [dbo].[{audit_table}]"),
)
.await
.expect("verify trigger rows");
let cleanup = format!(
"DROP TRIGGER IF EXISTS [dbo].[{trigger}]; DROP TABLE IF EXISTS [dbo].[{audit_table}]; DROP TABLE IF EXISTS [dbo].[{table}];"
);
let _ = dbx_core::db::sqlserver::execute_batch(&mut setup_client, &cleanup).await;
state.remove_connection_pools_detached(&connection_id).await;
let _ = std::fs::remove_dir_all(&dir);
assert_eq!(generated_summary.rows_imported, 1);
assert_eq!(explicit_summary.rows_imported, 1);
assert_eq!(plain_binary_summary.rows_imported, 1);
assert!(duplicate_error.is_err(), "unique constraint violation must fail the import");
assert_eq!(rows.rows.len(), 3);
assert_eq!(rows.rows[0][0], serde_json::json!("1"));
assert_eq!(rows.rows[0][1], serde_json::json!("auto"));
assert_eq!(rows.rows[0][2], serde_json::json!("2026-07-27T12:34:56.1234567"));
assert_eq!(rows.rows[0][3], serde_json::json!("12345678901234567890.1234567890"));
assert_eq!(rows.rows[0][4], serde_json::json!("\u{8d8a}\u{5357}\u{82b1}"));
assert_eq!(rows.rows[0][5], serde_json::json!("0x00ff10"));
assert_eq!(rows.rows[0][6], serde_json::json!("set"));
assert_eq!(rows.rows[1][0], serde_json::json!("42"));
assert_eq!(rows.rows[1][4], serde_json::json!("Tieng Viet"));
assert_eq!(rows.rows[1][5], serde_json::json!("0xabcdef"));
assert_eq!(rows.rows[2][1], serde_json::json!("plain"));
assert_eq!(rows.rows[2][7], serde_json::json!("plain"));
assert_eq!(audit_count.rows[0][0], serde_json::json!(3));
}
#[tokio::test]
#[ignore = "requires DBX_LIVE_SQLSERVER_HOST/PORT/USER/PASSWORD pointing at a writable SQL Server database"]
async fn live_sqlserver_table_import_bulk_cancels_before_target_write_and_rolls_back_truncate() {
let database = std::env::var("DBX_LIVE_SQLSERVER_DATABASE").unwrap_or_else(|_| "tempdb".to_string());
let suffix = uuid::Uuid::new_v4().simple().to_string();
let connection_id = format!("live-sqlserver-import-rollback-{suffix}");
let table = format!("dbx_bulk_rollback_{suffix}");
let mut setup_client = dbx_core::db::sqlserver::connect(
&std::env::var("DBX_LIVE_SQLSERVER_HOST").unwrap_or_else(|_| "127.0.0.1".to_string()),
std::env::var("DBX_LIVE_SQLSERVER_PORT").ok().and_then(|value| value.parse().ok()).unwrap_or(1433),
&std::env::var("DBX_LIVE_SQLSERVER_USER").unwrap_or_else(|_| "sa".to_string()),
&std::env::var("DBX_LIVE_SQLSERVER_PASSWORD").expect("DBX_LIVE_SQLSERVER_PASSWORD"),
Some(&database),
None,
Duration::from_secs(10),
)
.await
.expect("connect SQL Server");
dbx_core::db::sqlserver::execute_batch(
&mut setup_client,
&format!(
"CREATE TABLE [dbo].[{table}] ([id] INT NOT NULL PRIMARY KEY, [amount] DECIMAL(38,10) NOT NULL CHECK ([amount] > 0)); \
INSERT INTO [dbo].[{table}] VALUES (999, 9.0000000000);"
),
)
.await
.expect("create rollback target");
let (state, pool_key, dir) = live_sqlserver_import_state(&connection_id, &database, &suffix).await;
let invalid_csv = dir.join("invalid.csv");
std::fs::write(&invalid_csv, "id,amount\n1,-1.0000000000\n").expect("write invalid CSV");
let invalid_request = live_sqlserver_import_request(
&connection_id,
&database,
&table,
&invalid_csv,
vec![live_sqlserver_import_mapping("id", "id"), live_sqlserver_import_mapping("amount", "amount")],
TableImportMode::Truncate,
);
let truncate_error = run_live_sqlserver_import(&state, &pool_key, &invalid_request).await;
let rows_after_failed_truncate = dbx_core::db::sqlserver::execute_query(
&mut setup_client,
&format!("SELECT [id], CONVERT(VARCHAR(50), [amount]) FROM [dbo].[{table}]"),
)
.await
.expect("verify truncate rollback");
let valid_csv = dir.join("cancelled.csv");
std::fs::write(&valid_csv, "id,amount\n1,1.0000000000\n").expect("write cancellation CSV");
let cancel_request = live_sqlserver_import_request(
&connection_id,
&database,
&table,
&valid_csv,
vec![live_sqlserver_import_mapping("id", "id"), live_sqlserver_import_mapping("amount", "amount")],
TableImportMode::Truncate,
);
let cancellation_checks = Arc::new(AtomicUsize::new(0));
let cancellation_checks_for_import = cancellation_checks.clone();
let cancelled_progress = Arc::new(AtomicBool::new(false));
let cancelled_progress_for_import = cancelled_progress.clone();
let cancel_error = import_table_file_core(
&state,
&cancel_request,
&DatabaseType::SqlServer,
&pool_key,
move |_| {
let checks = cancellation_checks_for_import.clone();
Box::pin(async move { checks.fetch_add(1, Ordering::SeqCst) >= 1 })
},
move |progress| {
if progress.status == TableImportStatus::Cancelled {
cancelled_progress_for_import.store(true, Ordering::SeqCst);
}
},
)
.await;
let rows_after_cancel = dbx_core::db::sqlserver::execute_query(
&mut setup_client,
&format!("SELECT [id], CONVERT(VARCHAR(50), [amount]) FROM [dbo].[{table}]"),
)
.await
.expect("verify cancellation leaves target unchanged");
let _ = dbx_core::db::sqlserver::execute_batch(&mut setup_client, &format!("DROP TABLE IF EXISTS [dbo].[{table}]"))
.await;
state.remove_connection_pools_detached(&connection_id).await;
let _ = std::fs::remove_dir_all(&dir);
assert!(truncate_error.is_err(), "check constraint must fail the truncate import");
assert_eq!(rows_after_failed_truncate.rows, vec![vec![serde_json::json!(999), serde_json::json!("9.0000000000")]]);
assert_eq!(cancel_error.unwrap_err(), "Import cancelled");
assert!(cancellation_checks.load(Ordering::SeqCst) >= 2);
assert!(cancelled_progress.load(Ordering::SeqCst));
assert_eq!(rows_after_cancel.rows, vec![vec![serde_json::json!(999), serde_json::json!("9.0000000000")]]);
}
#[tokio::test]
#[ignore = "requires DBX_LIVE_SQLSERVER_HOST/PORT/USER/PASSWORD pointing at a writable SQL Server database"]
async fn live_sqlserver_column_metadata_marks_non_positional_insert_columns() {
+98
View File
@@ -0,0 +1,98 @@
# CSV/Excel 批量导入真实环境性能基准
**测试日期:2026 年 7 月 29 日**
本基准测试对比父提交 `a96cf1194` 的实现与本 PR 优化后的导入路径。测试覆盖完整的文件解析和数据库写入流程,而不只是 SQL 生成过程。
## 测试环境
- 客户端:Windows 11 64 位,Intel Core i7-13620H(10 核、16 个逻辑处理器),31.7 GiB 内存。
- 工具链:Rust 1.96.0,使用 release profile,`dbx-core` 通过 `--no-default-features` 构建。
- SQL Server:SQL Server 2022 `16.0.4265.3`,运行于本地 Docker 容器。
- PostgreSQL:PostgreSQL `16.14`,远程可写测试实例;未隔离网络波动及服务器负载。
- MySQL:MySQL `8.4.6`,远程可写测试实例,`max_allowed_packet = 64 MiB`;未隔离网络波动及服务器负载。
- 导入表结构:1 个 `BIGINT` 列和 11 个文本列。SQL Server、PostgreSQL 使用 `batch_size = 500`;MySQL 使用 `batch_size = 10000`,确保单个解析批次生成的 SQL 超过 512 KiB 目标并进入按字节拆批路径。
- CSV 数据集:40,889,006 字节,200,000 行 × 12 列。
- Excel 数据集:4,559,480 字节(压缩后的 XLSX),100,000 行 × 12 列,共 120 万个单元格。
每个场景运行 3 次,基线版本和当前版本交替执行,结果表采用中位数。文件生成和数据库初始化不计入计时,文件解析和数据库写入计入计时。吞吐测试期间每 10 ms 采样一次进程 RSS。取消测试在首次收到写入进度后发出取消请求,取消延迟指从发出请求到导入 Future 返回所需的时间。
测量前已校验独立构建的可执行文件:
- 基线版本 SHA-256:`6E1C95139F3360EDBD0BBC5D739FEC5DBD628911C288BA339424271B088077D9`
- 优化版本 SHA-256:`51D79649B2475BF5A0B40549A6F04FC880344881149F9BF7E2157A157183BB5C`
MySQL 补测在合并上游后重新独立构建,使用以下可执行文件:
- MySQL 基线版本 SHA-256:`A884D76A1A39BED0198AB72872E1CE96DB678ABB97E04737000E7465D8865868`
- MySQL 优化版本 SHA-256:`48418182C662CEC4B198D4AD0CF6376164B857E05B102526B269CB553EE7E4F3`
## 测试结果
| 数据库 | 数据源 | 导入路径 | 文件大小(字节) | 行数 × 列数 | 耗时(ms) | 吞吐量(行/秒) | 峰值 RSS(MiB) | RSS 增量(MiB) | 取消延迟(ms) |
| --- | --- | --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: |
| SQL Server | CSV | 生成 INSERT(`a96cf1194`) | 40,889,006 | 200,000 × 12 | 38,568.1 | 5,185.6 | 20.04 | 5.76 | 0.833 |
| SQL Server | CSV | TDS Bulk + NVARCHAR 暂存表 | 40,889,006 | 200,000 × 12 | 8,142.1 | 24,563.7 | 20.21 | 5.49 | 0.760 |
| SQL Server | XLSX | 生成 INSERT(`a96cf1194`) | 4,559,480 | 100,000 × 12 | 19,667.9 | 5,084.4 | 19.68 | 4.41 | 0.624 |
| SQL Server | XLSX | TDS Bulk + NVARCHAR 暂存表 | 4,559,480 | 100,000 × 12 | 5,271.6 | 18,969.5 | 19.62 | 4.56 | 0.482 |
| PostgreSQL | CSV | 每个解析批次执行一次 COPY(`a96cf1194`) | 40,889,006 | 200,000 × 12 | 36,243.0 | 5,518.3 | 23.03 | 5.70 | 1.536 |
| PostgreSQL | CSV | 8 MiB / 5 万行 COPY 累加器 | 40,889,006 | 200,000 × 12 | 4,967.9 | 40,258.7 | 37.82 | 20.29 | 1.202 |
| PostgreSQL | XLSX | 每个解析批次执行一次 COPY(`a96cf1194`) | 4,559,480 | 100,000 × 12 | 22,838.6 | 4,378.6 | 22.59 | 4.05 | 0.912 |
| PostgreSQL | XLSX | 8 MiB / 5 万行 COPY 累加器 | 4,559,480 | 100,000 × 12 | 4,160.7 | 24,034.2 | 37.09 | 18.93 | 0.855 |
| MySQL | CSV | 重复序列化并二分确定 512 KiB INSERT 批次(`a96cf1194`) | 40,889,006 | 200,000 × 12 | 33,786.3 | 5,919.6 | 100.33 | 84.79 | 13.659 |
| MySQL | CSV | 行值单次序列化 + SQL 字节自适应批次 | 40,889,006 | 200,000 × 12 | 18,793.6 | 10,641.9 | 93.55 | 79.20 | 12.746 |
| MySQL | XLSX | 重复序列化并二分确定 512 KiB INSERT 批次(`a96cf1194`) | 4,559,480 | 100,000 × 12 | 18,785.0 | 5,323.4 | 77.71 | 62.84 | 9.957 |
| MySQL | XLSX | 行值单次序列化 + SQL 字节自适应批次 | 4,559,480 | 100,000 × 12 | 11,210.4 | 8,920.3 | 71.69 | 56.43 | 9.489 |
吞吐量中位数变化如下:
- SQL Server CSV:提升至 `4.74x`(`+373.7%`)。
- SQL Server Excel:提升至 `3.73x`(`+273.1%`)。
- PostgreSQL CSV:提升至 `7.30x`(`+629.5%`)。
- PostgreSQL Excel:提升至 `5.49x`(`+448.9%`)。
- MySQL CSV:提升至 `1.80x`(`+79.8%`)。
- MySQL Excel:提升至 `1.68x`(`+67.6%`)。
PostgreSQL COPY 累加器以有界的内存开销换取更少的网络往返次数。在这些数据集上,其峰值 RSS 中位数增加约 14~15 MiB;内存使用仍受 COPY 累加器、编码后批次、解析器和驱动缓冲区共同约束。远程 PostgreSQL 的 CSV 测试存在网络波动,优化版本的吞吐量范围为 17,600~42,300 行/秒;表中报告的 40,258.7 行/秒是中位数,而非最好成绩。
MySQL 使用 10,000 行解析批次,使拆批前的候选 INSERT 约为 2 MiB,稳定超过 512 KiB SQL 目标。基线实现会反复序列化候选行并通过二分搜索确定拆分点;优化实现先将每行 SQL 值序列化一次,再按累计字节数线性拆批,因此本场景直接覆盖了被优化的路径。当前实现还会读取服务端 `max_allowed_packet` 并为单条 SQL 推导硬上限;本数据集没有触发 64 MiB 服务端包大小限制。大解析批次使 MySQL 两个版本的峰值 RSS 都明显高于 500 行场景,但优化版本的 CSV 和 XLSX 峰值 RSS 中位数分别降低 6.77 MiB 和 6.02 MiB。全部 MySQL 测试结束后再次查询测试 schema,`dbx_import_bench_%` 遗留表数量为 0。
## 边界场景补测
针对 SQL Server Bulk 转换内存评审,额外使用接近 TDS Bulk 单列编码上限的宽文本和最大有效批次运行 3 次。数据集为 180,034,911 字节 CSV、6,000 行 x 2 列,每个文本值 30,000 字节。命令请求 `batch_size = 5000`,但 SQL Server 导入会将其限制为 1,000 行,因此每个实际解析批次包含约 28.6 MiB 原始文本。结果仍取中位数:
| 场景 | 耗时(ms) | 吞吐(行/秒) | 峰值 RSS(MiB) | RSS 增量(MiB) | 取消延迟(ms) |
| --- | ---: | ---: | ---: | ---: | ---: |
| SQL Server 宽字段大批次 | 5,663.7 | 1,059.4 | 132.30 | 117.91 | 20.866 |
该场景覆盖 SQL Server 允许的最大有效批次。解析线程、两槽有界通道和数据库消费者可能同时持有多个原始数据批次,因此进程 RSS 包含这些有界的源数据;Bulk 转换本身使用逐行转换和逐行 TDS 发送,不再创建批次级完整字符串矩阵,转换后的附加内存受 16 MiB 单行预算约束。单元回归同时覆盖 32 行 x 每行 1 MiB 的惰性转换,以及单行超过 16 MiB 预算时在复制前拒绝。Tiberius 当前会拒绝 UTF-16 编码长度超过 65,535 字节的单个 Bulk 字符串,因此真实写入补测使用 30,000 字节 ASCII 文本;1 MiB 单列输入会在驱动编码层被拒绝,不能作为成功写入基准。
针对 XLSX 首次写入前取消,回归夹具包含 8,193 个共享字符串,`sharedStrings.xml` 正文约 4.16 MiB。3 次取消计时为 162.64 ms、139.22 ms 和 153.67 ms,中位数为 153.67 ms;取消均发生在 Header 和任何数据库写入之前。读取器每 64 KiB 检查共享取消状态,异步侧每 25 ms 轮询一次,预校验和正式解析分别占文件读取进度的前、后 50%,进度保持单调。
宽字段场景可使用以下附加参数复现:
```powershell
cargo run -p dbx-core --no-default-features --release `
--example table_import_live_bench -- `
--database=sqlserver --format=csv --rows=6000 --columns=2 `
--batch-size=5000 --text-bytes=30000
```
## 复现方式
数据库连接信息仅通过环境变量提供:
```powershell
$env:DBX_BENCH_HOST = '<host>'
$env:DBX_BENCH_PORT = '<port>'
$env:DBX_BENCH_USER = '<user>'
$env:DBX_BENCH_PASSWORD = '<password>'
$env:DBX_BENCH_DATABASE = '<database>'
$env:DBX_BENCH_SCHEMA = '<schema>'
cargo run -p dbx-core --no-default-features --release `
--example table_import_live_bench -- `
--database=postgres --format=csv --rows=200000 --columns=12 --batch-size=500
```
测试 MySQL 或 SQL Server 时分别使用 `--database=mysql` 或 `--database=sqlserver`;MySQL 性能补测同时使用 `--batch-size=10000`。测试 Excel 场景时使用 `--format=xlsx --rows=100000`。基准程序会创建名称唯一的临时表,依次执行吞吐测试和取消测试,输出一条 JSON 结果,并在结束时删除测试表和临时文件。