diff --git a/apps/desktop/src/components/import/TableImportDialog.vue b/apps/desktop/src/components/import/TableImportDialog.vue index cc384afc3..4133a005f 100644 --- a/apps/desktop/src/components/import/TableImportDialog.vue +++ b/apps/desktop/src/components/import/TableImportDialog.vue @@ -8,10 +8,13 @@ import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; import { Label } from "@/components/ui/label"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { SearchableSelect } from "@/components/ui/searchable-select"; import { AlertTriangle, ArrowLeft, ArrowRight, Check, CheckCircle2, FileJson, FileSpreadsheet, FileText, FileUp, Loader2, RefreshCw, Square, Upload, X } from "@lucide/vue"; import { useConnectionStore } from "@/stores/connectionStore"; import { useToast } from "@/composables/useToast"; -import { autoMapImportColumns, nextTableImportWizardStep, previousTableImportWizardStep, requiredImportTargetColumns, validateImportMappings, type TableImportWizardStep } from "@/lib/table/tableImport"; +import { autoMapImportColumns, nextTableImportWizardStep, previousTableImportWizardStep, requiredImportTargetColumns, suggestImportTargetDataTypes, validateImportMappings, type TableImportWizardStep } from "@/lib/table/tableImport"; +import { getDataTypeOptions } from "@/lib/table/tableStructureEditorState"; +import { tableStructureDatabaseTypeForConnection } from "@/lib/database/jdbcDialect"; import type { ColumnInfo } from "@/types/database"; import * as api from "@/lib/backend/api"; @@ -37,6 +40,9 @@ const selectedSource = ref(null); const sourceFormat = ref("csv"); const preview = ref(null); const columnMapping = ref>({}); +const columnDataTypes = ref>({}); +const dynamicDataTypeOptions = ref([]); +const loadingDataTypeOptions = ref(false); const loadingTarget = ref(false); const loadingPreview = ref(false); const importMode = ref("append"); @@ -56,6 +62,7 @@ const selectedSheet = ref(""); const jsonShape = ref("auto"); const previewLimit = ref(50); let previewReloadTimer: ReturnType | null = null; +let dataTypeOptionsRequestId = 0; const formatOptions: Array<{ value: api.TableImportSourceFormat; icon: any; labelKey: string; descriptionKey: string }> = [ { value: "csv", icon: FileText, labelKey: "tableImport.formatCsv", descriptionKey: "tableImport.formatCsvDescription" }, @@ -74,6 +81,8 @@ const wizardSteps: Array<{ value: TableImportWizardStep; labelKey: string }> = [ ]; const selectedConnection = computed(() => (props.prefillConnectionId ? store.getConfig(props.prefillConnectionId) : undefined)); +const structureDatabaseType = computed(() => tableStructureDatabaseTypeForConnection(selectedConnection.value)); +const dataTypeOptions = computed(() => mergeDataTypeOptions(dynamicDataTypeOptions.value, getDataTypeOptions(structureDatabaseType.value), Object.values(columnDataTypes.value))); const hasExistingTarget = computed(() => !!props.prefillTable); const targetTableName = computed(() => (targetMode.value === "create" ? newTableName.value.trim() : props.prefillTable || "")); const targetColumnNames = computed(() => targetColumns.value.map((column) => column.name)); @@ -81,10 +90,14 @@ const mappedColumns = computed(() => { const currentPreview = preview.value; if (!currentPreview) return []; return currentPreview.columns - .map((sourceColumn) => ({ - sourceColumn, - targetColumn: columnMapping.value[sourceColumn] ?? "", - })) + .map((sourceColumn) => { + const targetDataType = targetMode.value === "create" ? String(columnDataTypes.value[sourceColumn] ?? "").trim() : undefined; + return { + sourceColumn, + targetColumn: columnMapping.value[sourceColumn] ?? "", + ...(targetMode.value === "create" ? { targetDataType } : {}), + }; + }) .filter((mapping) => mapping.targetColumn); }); const mappedCount = computed(() => mappedColumns.value.length); @@ -118,6 +131,13 @@ const selectedSourceName = computed(() => { if (!source) return ""; return typeof source === "string" ? source.split(/[\\/]/).pop() || source : source.name; }); +const createColumnSummaries = computed(() => + mappedColumns.value.map((mapping) => ({ + sourceColumn: mapping.sourceColumn, + targetColumn: mapping.targetColumn, + targetDataType: mapping.targetDataType || "", + })), +); const parseOptions = computed(() => ({ delimiter: sourceFormat.value === "tsv" ? "\\t" : sourceFormat.value === "csv" ? "," : delimiter.value, hasHeader: hasHeader.value, @@ -143,6 +163,7 @@ function resetState() { previewLimit.value = 50; preview.value = null; columnMapping.value = {}; + columnDataTypes.value = {}; importMode.value = "append"; batchSize.value = 500; running.value = false; @@ -168,6 +189,22 @@ function suggestedTableName(name: string) { return withoutExtension.replace(/[\s-]+/g, "_") || "imported_data"; } +function mergeDataTypeOptions(...groups: readonly string[][]): string[] { + const seen = new Set(); + const result: string[] = []; + for (const group of groups) { + for (const option of group) { + const trimmed = option.trim(); + if (!trimmed) continue; + const key = trimmed.toLowerCase(); + if (seen.has(key)) continue; + seen.add(key); + result.push(trimmed); + } + } + return result; +} + function applyAutoMapping() { const currentPreview = preview.value; if (!currentPreview) return; @@ -178,6 +215,42 @@ function applyAutoMapping() { columnMapping.value = autoMapImportColumns(currentPreview.columns, targetColumnNames.value); } +function applySuggestedColumnDataTypes(currentPreview = preview.value) { + if (targetMode.value !== "create" || !currentPreview) { + columnDataTypes.value = {}; + return; + } + const suggested = suggestImportTargetDataTypes(currentPreview.columns, currentPreview.rows, structureDatabaseType.value); + const previous = columnDataTypes.value; + columnDataTypes.value = Object.fromEntries(currentPreview.columns.map((sourceColumn) => [sourceColumn, previous[sourceColumn]?.trim() ? previous[sourceColumn] : suggested[sourceColumn] || "TEXT"])); +} + +async function loadDataTypeOptions() { + const requestId = ++dataTypeOptionsRequestId; + const connectionId = props.prefillConnectionId; + const database = props.prefillDatabase || ""; + if (!connectionId || !database || targetMode.value !== "create") { + dynamicDataTypeOptions.value = []; + loadingDataTypeOptions.value = false; + return; + } + loadingDataTypeOptions.value = true; + try { + await store.ensureConnected(connectionId); + const options = await api.listDataTypes(connectionId, database); + if (requestId !== dataTypeOptionsRequestId) return; + dynamicDataTypeOptions.value = mergeDataTypeOptions(options); + } catch { + if (requestId === dataTypeOptionsRequestId) { + dynamicDataTypeOptions.value = []; + } + } finally { + if (requestId === dataTypeOptionsRequestId) { + loadingDataTypeOptions.value = false; + } + } +} + async function loadTargetColumns() { if (targetMode.value !== "existing" || !props.prefillConnectionId || !props.prefillDatabase || !props.prefillTable) return; loadingTarget.value = true; @@ -212,9 +285,11 @@ async function loadPreview(fileOrPath = selectedSource.value) { selectedSheet.value = nextPreview.sheets[0]; } applyAutoMapping(); + applySuggestedColumnDataTypes(nextPreview); } catch (e: any) { preview.value = null; columnMapping.value = {}; + columnDataTypes.value = {}; errorMessage.value = String(e?.message || e); } finally { loadingPreview.value = false; @@ -225,6 +300,7 @@ function assignSelectedSource(source: string | File) { selectedSource.value = source; preview.value = null; columnMapping.value = {}; + columnDataTypes.value = {}; progress.value = null; errorMessage.value = ""; const name = typeof source === "string" ? source : source.name; @@ -272,6 +348,13 @@ function updateMapping(sourceColumn: string, value: any) { }; } +function updateColumnDataType(sourceColumn: string, value: any) { + columnDataTypes.value = { + ...columnDataTypes.value, + [sourceColumn]: String(value), + }; +} + function formatCell(value: unknown) { if (value === null) return "NULL"; if (typeof value === "object") return JSON.stringify(value); @@ -408,6 +491,7 @@ watch( if (value) { resetState(); void loadTargetColumns(); + void loadDataTypeOptions(); } }, { immediate: true }, @@ -416,11 +500,15 @@ watch( watch([sourceFormat, delimiter, hasHeader, trimValues, emptyStringAsNull, selectedSheet, jsonShape, previewLimit], schedulePreviewReload); watch(targetMode, (mode) => { if (mode === "existing") { + columnDataTypes.value = {}; + dynamicDataTypeOptions.value = []; void loadTargetColumns(); } else { targetColumns.value = []; importMode.value = "append"; applyAutoMapping(); + applySuggestedColumnDataTypes(); + void loadDataTypeOptions(); } }); @@ -621,11 +709,16 @@ watch(targetMode, (mode) => { -
+
{{ t("tableImport.mapping") }}
-
+
+ {{ t("tableImport.sourceColumn") }} + {{ t("tableImport.targetColumn") }} + {{ t("tableImport.targetDataType") }} +
+
{{ sourceColumn }}
@@ -641,6 +734,22 @@ watch(targetMode, (mode) => { +
@@ -714,6 +823,33 @@ watch(targetMode, (mode) => {
+
+
{{ t("tableImport.createColumns") }}
+
+ + + + + + + + + + + + + + + +
{{ t("tableImport.sourceColumn") }}{{ t("tableImport.targetColumn") }}{{ t("tableImport.targetDataType") }}
+ {{ column.sourceColumn }} + + {{ column.targetColumn }} + + {{ column.targetDataType }} +
+
+
{{ t("tableImport.truncateWarning") }} diff --git a/apps/desktop/src/i18n/locales/en.ts b/apps/desktop/src/i18n/locales/en.ts index 19ff9a2a2..86f82255f 100644 --- a/apps/desktop/src/i18n/locales/en.ts +++ b/apps/desktop/src/i18n/locales/en.ts @@ -2283,6 +2283,10 @@ export default { rows: "Rows", mapped: "Mapped", mapping: "Column Mapping", + sourceColumn: "Source", + targetColumn: "Target", + targetDataType: "Type", + createColumns: "Columns to create", skipColumn: "Skip", preview: "Preview", mode: "Import mode", diff --git a/apps/desktop/src/i18n/locales/es.ts b/apps/desktop/src/i18n/locales/es.ts index ce89a3b36..b9cd17b0a 100644 --- a/apps/desktop/src/i18n/locales/es.ts +++ b/apps/desktop/src/i18n/locales/es.ts @@ -2219,6 +2219,10 @@ export default withEnglishFallback({ rows: "Filas", mapped: "Mapeadas", mapping: "Mapeo de columnas", + sourceColumn: "Origen", + targetColumn: "Destino", + targetDataType: "Tipo", + createColumns: "Columnas a crear", skipColumn: "Omitir", preview: "Vista previa", mode: "Modo de importación", diff --git a/apps/desktop/src/i18n/locales/it.ts b/apps/desktop/src/i18n/locales/it.ts index a0884cd0b..ddafb0720 100644 --- a/apps/desktop/src/i18n/locales/it.ts +++ b/apps/desktop/src/i18n/locales/it.ts @@ -2217,6 +2217,10 @@ export default withEnglishFallback({ rows: "Righe", mapped: "Mappata", mapping: "Mappatura Colonne", + sourceColumn: "Origine", + targetColumn: "Destinazione", + targetDataType: "Tipo", + createColumns: "Colonne da creare", skipColumn: "Salta", preview: "Anteprima", mode: "Modalità importazione", diff --git a/apps/desktop/src/i18n/locales/ja.ts b/apps/desktop/src/i18n/locales/ja.ts index a4303161f..c4402fa1f 100644 --- a/apps/desktop/src/i18n/locales/ja.ts +++ b/apps/desktop/src/i18n/locales/ja.ts @@ -2217,6 +2217,10 @@ export default withEnglishFallback({ rows: "行数", mapped: "マッピング済み", mapping: "列マッピング", + sourceColumn: "ソース", + targetColumn: "ターゲット", + targetDataType: "型", + createColumns: "作成する列", skipColumn: "スキップ", preview: "プレビュー", mode: "インポートモード", diff --git a/apps/desktop/src/i18n/locales/pt-BR.ts b/apps/desktop/src/i18n/locales/pt-BR.ts index cdc0ae770..466850a35 100644 --- a/apps/desktop/src/i18n/locales/pt-BR.ts +++ b/apps/desktop/src/i18n/locales/pt-BR.ts @@ -2218,6 +2218,10 @@ export default withEnglishFallback({ rows: "Linhas", mapped: "Mapeado", mapping: "Mapeamento de Colunas", + sourceColumn: "Origem", + targetColumn: "Destino", + targetDataType: "Tipo", + createColumns: "Colunas a criar", skipColumn: "Pular", preview: "Pré-visualização", mode: "Modo de importação", diff --git a/apps/desktop/src/i18n/locales/zh-CN.ts b/apps/desktop/src/i18n/locales/zh-CN.ts index 61750a9b1..f8499d0a5 100644 --- a/apps/desktop/src/i18n/locales/zh-CN.ts +++ b/apps/desktop/src/i18n/locales/zh-CN.ts @@ -2283,6 +2283,10 @@ export default withEnglishFallback({ rows: "行数", mapped: "已映射", mapping: "字段映射", + sourceColumn: "来源字段", + targetColumn: "目标字段", + targetDataType: "目标类型", + createColumns: "将创建的字段", skipColumn: "跳过", preview: "预览", mode: "导入模式", diff --git a/apps/desktop/src/i18n/locales/zh-TW.ts b/apps/desktop/src/i18n/locales/zh-TW.ts index 888945ff3..76a63a273 100644 --- a/apps/desktop/src/i18n/locales/zh-TW.ts +++ b/apps/desktop/src/i18n/locales/zh-TW.ts @@ -2120,6 +2120,10 @@ export default withEnglishFallback({ rows: "列數", mapped: "已映射", mapping: "欄位映射", + sourceColumn: "來源欄位", + targetColumn: "目標欄位", + targetDataType: "目標類型", + createColumns: "將建立的欄位", skipColumn: "跳過", preview: "預覽", mode: "匯入模式", diff --git a/apps/desktop/src/lib/__tests__/table/tableImport.spec.ts b/apps/desktop/src/lib/__tests__/table/tableImport.spec.ts index 2099ee94f..6887130b9 100644 --- a/apps/desktop/src/lib/__tests__/table/tableImport.spec.ts +++ b/apps/desktop/src/lib/__tests__/table/tableImport.spec.ts @@ -1,5 +1,5 @@ import { describe, expect, it } from "vitest"; -import { autoMapImportColumns, nextTableImportWizardStep, previousTableImportWizardStep, requiredImportTargetColumns, validateImportMappings } from "@/lib/table/tableImport"; +import { autoMapImportColumns, nextTableImportWizardStep, previousTableImportWizardStep, requiredImportTargetColumns, suggestImportTargetDataTypes, validateImportMappings } from "@/lib/table/tableImport"; describe("tableImport", () => { it("auto maps exact and normalized column names", () => { @@ -27,6 +27,13 @@ describe("tableImport", () => { expect(result.errors[0]).toContain("Target column mapped more than once"); }); + it("rejects empty create-table data types", () => { + const result = validateImportMappings([{ sourceColumn: "code", targetColumn: "code", targetDataType: "" }]); + + expect(result.valid).toBe(false); + expect(result.errors).toEqual(["Target data type cannot be empty: code"]); + }); + it("detects unmapped required target columns", () => { expect( requiredImportTargetColumns( @@ -46,4 +53,22 @@ describe("tableImport", () => { expect(previousTableImportWizardStep("review")).toBe("mapping"); expect(previousTableImportWizardStep("source")).toBe("source"); }); + + it("suggests create-table data types from preview rows", () => { + expect( + suggestImportTargetDataTypes( + ["id", "code", "amount", "created_at"], + [ + ["1001", "00123", "12.5", "2026-07-07 08:15:00"], + ["1002", "00456", "13.75", "2026-07-07 09:15:00"], + ], + "mysql", + ), + ).toEqual({ + id: "BIGINT", + code: "TEXT", + amount: "DOUBLE", + created_at: "DATETIME", + }); + }); }); diff --git a/apps/desktop/src/lib/backend/tauri.ts b/apps/desktop/src/lib/backend/tauri.ts index 41ad218a0..cb61525ef 100644 --- a/apps/desktop/src/lib/backend/tauri.ts +++ b/apps/desktop/src/lib/backend/tauri.ts @@ -1921,6 +1921,7 @@ export type TableImportJsonShape = "auto" | "objects" | "arrays"; export interface TableImportColumnMapping { sourceColumn: string; targetColumn: string; + targetDataType?: string | null; } export interface TableImportParseOptions { diff --git a/apps/desktop/src/lib/table/tableImport.ts b/apps/desktop/src/lib/table/tableImport.ts index 5b93d3317..13b04c3f5 100644 --- a/apps/desktop/src/lib/table/tableImport.ts +++ b/apps/desktop/src/lib/table/tableImport.ts @@ -1,8 +1,11 @@ +import type { DatabaseType } from "@/types/database"; + export const IMPORT_SKIP_TARGET = ""; export interface ImportColumnMappingLike { sourceColumn: string; targetColumn: string; + targetDataType?: string | null; } export interface ImportMappingValidationResult { @@ -49,6 +52,9 @@ export function validateImportMappings(mappings: ImportColumnMappingLike[]): Imp duplicateTargets.push(mapping.targetColumn); } seen.add(key); + if (Object.prototype.hasOwnProperty.call(mapping, "targetDataType") && !String(mapping.targetDataType || "").trim()) { + errors.push(`Target data type cannot be empty: ${mapping.targetColumn}`); + } } if (duplicateTargets.length) { errors.push(`Target column mapped more than once: ${duplicateTargets.join(", ")}`); @@ -81,3 +87,109 @@ export function previousTableImportWizardStep(step: TableImportWizardStep): Tabl const index = TABLE_IMPORT_WIZARD_STEPS.indexOf(step); return TABLE_IMPORT_WIZARD_STEPS[Math.max(0, index - 1)]; } + +type ImportInferredType = "boolean" | "integer" | "decimal" | "date" | "timestamp" | "json" | "text"; + +function hasNumericLeadingZero(value: string): boolean { + const unsigned = value.trim().replace(/^[+-]/, ""); + return unsigned.length > 1 && unsigned[0] === "0" && /\d/.test(unsigned[1] || ""); +} + +function isLikelyDate(value: string): boolean { + return /^\d{4}[-/]\d{2}[-/]\d{2}$/.test(value.trim()); +} + +function isLikelyTimestamp(value: string): boolean { + const trimmed = value.trim(); + return /^\d{4}[-/]\d{2}[-/]\d{2}[ T]\d{2}:\d{2}:\d{2}(?:\.\d+)?$/.test(trimmed) || /^\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:\d{2})$/.test(trimmed); +} + +function inferStringType(value: string): ImportInferredType { + const trimmed = value.trim(); + if (!trimmed) return "text"; + if (isLikelyTimestamp(trimmed)) return "timestamp"; + if (isLikelyDate(trimmed)) return "date"; + if (!hasNumericLeadingZero(trimmed)) { + if (/^[+-]?\d+$/.test(trimmed)) return "integer"; + if (/^[+-]?(?:\d+\.\d*|\d*\.\d+|\d+e[+-]?\d+|\d+\.\d*e[+-]?\d+|\d*\.\d+e[+-]?\d+)$/i.test(trimmed) && Number.isFinite(Number(trimmed))) { + return "decimal"; + } + } + return "text"; +} + +function inferValueType(value: unknown): ImportInferredType | null { + if (value == null) return null; + if (typeof value === "boolean") return "boolean"; + if (typeof value === "number") return Number.isInteger(value) ? "integer" : "decimal"; + if (typeof value === "string") return inferStringType(value); + if (typeof value === "object") return "json"; + return "text"; +} + +function mergeInferredType(current: ImportInferredType | null, next: ImportInferredType): ImportInferredType { + if (!current || current === next) return next; + if (current === "text" || next === "text") return "text"; + if ((current === "integer" && next === "decimal") || (current === "decimal" && next === "integer")) return "decimal"; + if ((current === "date" && next === "timestamp") || (current === "timestamp" && next === "date")) return "timestamp"; + return "text"; +} + +function inferColumnType(rows: unknown[][], sourceIndex: number): ImportInferredType { + let inferred: ImportInferredType | null = null; + for (const row of rows) { + const valueType = inferValueType(row[sourceIndex]); + if (!valueType) continue; + inferred = mergeInferredType(inferred, valueType); + if (inferred === "text") break; + } + return inferred || "text"; +} + +export function importDataTypeForDatabase(inferredType: ImportInferredType, databaseType?: DatabaseType): string { + switch (inferredType) { + case "boolean": + if (["mysql", "doris", "starrocks", "goldendb", "sundb", "databend"].includes(databaseType || "")) return "TINYINT(1)"; + if (databaseType === "sqlserver") return "BIT"; + if (databaseType === "sqlite" || databaseType === "rqlite" || databaseType === "turso") return "INTEGER"; + if (databaseType === "oracle" || databaseType === "oceanbase-oracle" || databaseType === "dameng") return "NUMBER(1)"; + if (databaseType === "clickhouse") return "UInt8"; + return "BOOLEAN"; + case "integer": + if (databaseType === "sqlite" || databaseType === "rqlite" || databaseType === "turso") return "INTEGER"; + if (databaseType === "oracle" || databaseType === "oceanbase-oracle" || databaseType === "dameng") return "NUMBER(19)"; + if (databaseType === "clickhouse") return "Int64"; + return "BIGINT"; + case "decimal": + if (["postgres", "gaussdb", "opengauss", "redshift", "kingbase", "highgo", "kwdb", "vastbase"].includes(databaseType || "")) return "DOUBLE PRECISION"; + if (databaseType === "sqlite" || databaseType === "rqlite" || databaseType === "turso") return "REAL"; + if (databaseType === "oracle" || databaseType === "oceanbase-oracle" || databaseType === "dameng") return "BINARY_DOUBLE"; + if (databaseType === "clickhouse") return "Float64"; + return "DOUBLE"; + case "date": + if (databaseType === "sqlite" || databaseType === "rqlite" || databaseType === "turso") return "TEXT"; + if (databaseType === "clickhouse") return "Date"; + return "DATE"; + case "timestamp": + if (["mysql", "doris", "starrocks", "goldendb", "sundb", "databend"].includes(databaseType || "")) return "DATETIME"; + if (databaseType === "sqlserver") return "DATETIME2"; + if (databaseType === "sqlite" || databaseType === "rqlite" || databaseType === "turso") return "TEXT"; + if (databaseType === "clickhouse") return "DateTime64"; + return "TIMESTAMP"; + case "json": + if (["postgres", "gaussdb", "opengauss", "kingbase", "highgo", "kwdb", "vastbase"].includes(databaseType || "")) return "JSONB"; + if (databaseType === "mysql" || databaseType === "databend") return "JSON"; + return importDataTypeForDatabase("text", databaseType); + case "text": + default: + if (databaseType === "sqlserver") return "NVARCHAR(MAX)"; + if (databaseType === "oracle" || databaseType === "oceanbase-oracle" || databaseType === "dameng") return "CLOB"; + if (databaseType === "clickhouse") return "String"; + if (["hive", "trino", "prestosql", "databricks"].includes(databaseType || "")) return "STRING"; + return "TEXT"; + } +} + +export function suggestImportTargetDataTypes(columns: string[], rows: unknown[][], databaseType?: DatabaseType): Record { + return Object.fromEntries(columns.map((column, index) => [column, importDataTypeForDatabase(inferColumnType(rows, index), databaseType)])); +} diff --git a/crates/dbx-core/src/table_import.rs b/crates/dbx-core/src/table_import.rs index 49b0e65ff..318afbc64 100644 --- a/crates/dbx-core/src/table_import.rs +++ b/crates/dbx-core/src/table_import.rs @@ -47,6 +47,8 @@ pub struct ImportCreateTablePlan { pub struct TableImportColumnMapping { pub source_column: String, pub target_column: String, + #[serde(default)] + pub target_data_type: Option, } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -696,6 +698,15 @@ pub fn mapping_indexes_for_columns( columns: &[String], mappings: &[TableImportColumnMapping], ) -> Result, String> { + mapping_indexes_with_mappings(columns, mappings).map(|mapped| { + mapped.into_iter().map(|(source_index, mapping)| (source_index, mapping.target_column.clone())).collect() + }) +} + +fn mapping_indexes_with_mappings<'a>( + columns: &[String], + mappings: &'a [TableImportColumnMapping], +) -> Result, String> { if mappings.is_empty() { return Err("No columns mapped for import".to_string()); } @@ -712,7 +723,7 @@ pub fn mapping_indexes_for_columns( if !target_seen.insert(mapping.target_column.clone()) { return Err(format!("Target column mapped more than once: {}", mapping.target_column)); } - mapped.push((source_index, mapping.target_column.clone())); + mapped.push((source_index, mapping)); } Ok(mapped) } @@ -1011,6 +1022,52 @@ fn import_data_type(inferred_type: ImportInferredType, db_type: &DatabaseType) - .to_string() } +fn normalize_import_target_data_type(mapping: &TableImportColumnMapping) -> Result, String> { + let Some(raw_data_type) = mapping.target_data_type.as_deref() else { + return Ok(None); + }; + let data_type = raw_data_type.trim(); + if data_type.is_empty() { + return Err(format!("Target data type cannot be empty: {}", mapping.target_column)); + } + validate_import_target_data_type(data_type)?; + Ok(Some(data_type.to_string())) +} + +fn validate_import_target_data_type(data_type: &str) -> Result<(), String> { + let lowered = data_type.to_ascii_lowercase(); + if data_type.contains(';') + || lowered.contains("--") + || lowered.contains("/*") + || lowered.contains("*/") + || data_type.chars().any(char::is_control) + { + return Err(format!("Unsupported target data type syntax: {data_type}")); + } + + // A user-entered type is a DDL fragment, so keep it constrained to one type + // expression and reject separators that could add another column or clause. + let mut paren_depth = 0usize; + for ch in data_type.chars() { + match ch { + '(' => paren_depth += 1, + ')' => { + paren_depth = paren_depth + .checked_sub(1) + .ok_or_else(|| format!("Unsupported target data type syntax: {data_type}"))?; + } + ',' if paren_depth == 0 => { + return Err(format!("Unsupported target data type syntax: {data_type}")); + } + _ => {} + } + } + if paren_depth != 0 { + return Err(format!("Unsupported target data type syntax: {data_type}")); + } + Ok(()) +} + pub fn build_import_create_table_plan( data: &ParsedImportFile, mappings: &[TableImportColumnMapping], @@ -1021,12 +1078,17 @@ pub fn build_import_create_table_plan( if table.trim().is_empty() { return Err("Target table name is required".to_string()); } - let mapped = mapping_indexes(data, mappings)?; + let mapped = mapping_indexes_with_mappings(&data.columns, mappings)?; let mut columns = Vec::with_capacity(mapped.len()); - for (source_index, target_column) in mapped { - let inferred_type = infer_column_type(&data.rows, source_index); - columns - .push(ImportCreateTableColumn { name: target_column, data_type: import_data_type(inferred_type, db_type) }); + for (source_index, mapping) in mapped { + let data_type = match normalize_import_target_data_type(mapping)? { + Some(data_type) => data_type, + None => { + let inferred_type = infer_column_type(&data.rows, source_index); + import_data_type(inferred_type, db_type) + } + }; + columns.push(ImportCreateTableColumn { name: mapping.target_column.clone(), data_type }); } if columns.is_empty() { return Err("No columns mapped for import".to_string()); @@ -1694,7 +1756,11 @@ mod tests { let mappings = data .columns .iter() - .map(|column| TableImportColumnMapping { source_column: column.clone(), target_column: column.clone() }) + .map(|column| TableImportColumnMapping { + source_column: column.clone(), + target_column: column.clone(), + target_data_type: None, + }) .collect::>(); let plan = @@ -1721,8 +1787,11 @@ mod tests { fn create_table_plan_requires_target_table_name() { let data = ParsedImportFile { columns: vec!["id".to_string()], rows: vec![vec![serde_json::json!(1)]], total_rows: 1 }; - let mappings = - vec![TableImportColumnMapping { source_column: "id".to_string(), target_column: "id".to_string() }]; + let mappings = vec![TableImportColumnMapping { + source_column: "id".to_string(), + target_column: "id".to_string(), + target_data_type: None, + }]; let error = build_import_create_table_plan(&data, &mappings, " ", "", &DatabaseType::Mysql).unwrap_err(); @@ -1736,19 +1805,80 @@ mod tests { rows: vec![vec![serde_json::json!("long text")]], total_rows: 1, }; - let mappings = - vec![TableImportColumnMapping { source_column: "notes".to_string(), target_column: "notes".to_string() }]; + let mappings = vec![TableImportColumnMapping { + source_column: "notes".to_string(), + target_column: "notes".to_string(), + target_data_type: None, + }]; let plan = build_import_create_table_plan(&data, &mappings, "events", "dbo", &DatabaseType::SqlServer).unwrap(); assert_eq!(plan.sql, "CREATE TABLE [dbo].[events] (\n [notes] NVARCHAR(MAX)\n)"); } + #[test] + fn create_table_plan_uses_user_defined_column_type() { + let data = ParsedImportFile { + columns: vec!["code".to_string(), "amount".to_string()], + rows: vec![vec![serde_json::json!("1001"), serde_json::json!("12.5")]], + total_rows: 1, + }; + let mappings = vec![ + TableImportColumnMapping { + source_column: "code".to_string(), + target_column: "code".to_string(), + target_data_type: Some("VARCHAR(32)".to_string()), + }, + TableImportColumnMapping { + source_column: "amount".to_string(), + target_column: "amount".to_string(), + target_data_type: Some("DECIMAL(10,2)".to_string()), + }, + ]; + + let plan = build_import_create_table_plan(&data, &mappings, "invoice", "", &DatabaseType::Mysql).unwrap(); + + assert_eq!(plan.sql, "CREATE TABLE `invoice` (\n `code` VARCHAR(32),\n `amount` DECIMAL(10,2)\n)"); + assert_eq!( + plan.columns, + vec![ + ImportCreateTableColumn { name: "code".to_string(), data_type: "VARCHAR(32)".to_string() }, + ImportCreateTableColumn { name: "amount".to_string(), data_type: "DECIMAL(10,2)".to_string() }, + ] + ); + } + + #[test] + fn create_table_plan_rejects_unsafe_user_defined_column_type() { + let data = ParsedImportFile { + columns: vec!["name".to_string()], + rows: vec![vec![serde_json::json!("Ada")]], + total_rows: 1, + }; + let mappings = vec![TableImportColumnMapping { + source_column: "name".to_string(), + target_column: "name".to_string(), + target_data_type: Some("TEXT, injected INT".to_string()), + }]; + + let error = build_import_create_table_plan(&data, &mappings, "users", "", &DatabaseType::Mysql).unwrap_err(); + + assert!(error.contains("Unsupported target data type syntax")); + } + #[test] fn builds_import_insert_batches_from_mapped_columns() { let mappings = vec![ - TableImportColumnMapping { source_column: "id".to_string(), target_column: "user_id".to_string() }, - TableImportColumnMapping { source_column: "name".to_string(), target_column: "display_name".to_string() }, + TableImportColumnMapping { + source_column: "id".to_string(), + target_column: "user_id".to_string(), + target_data_type: None, + }, + TableImportColumnMapping { + source_column: "name".to_string(), + target_column: "display_name".to_string(), + target_data_type: None, + }, ]; let data = ParsedImportFile { columns: vec!["id".to_string(), "name".to_string(), "ignored".to_string()], @@ -1779,8 +1909,16 @@ mod tests { fn duplicate_mapping_is_rejected_before_sql_generation() { let columns = vec!["id".to_string(), "name".to_string()]; let mappings = vec![ - TableImportColumnMapping { source_column: "id".to_string(), target_column: "target".to_string() }, - TableImportColumnMapping { source_column: "name".to_string(), target_column: "target".to_string() }, + TableImportColumnMapping { + source_column: "id".to_string(), + target_column: "target".to_string(), + target_data_type: None, + }, + TableImportColumnMapping { + source_column: "name".to_string(), + target_column: "target".to_string(), + target_data_type: None, + }, ]; let error = mapping_indexes_for_columns(&columns, &mappings).unwrap_err(); @@ -1792,8 +1930,16 @@ mod tests { fn builds_single_streaming_import_batch_from_rows() { let columns = vec!["id".to_string(), "name".to_string()]; let mappings = vec![ - TableImportColumnMapping { source_column: "id".to_string(), target_column: "id".to_string() }, - TableImportColumnMapping { source_column: "name".to_string(), target_column: "name".to_string() }, + TableImportColumnMapping { + source_column: "id".to_string(), + target_column: "id".to_string(), + target_data_type: None, + }, + TableImportColumnMapping { + source_column: "name".to_string(), + target_column: "name".to_string(), + target_data_type: None, + }, ]; let rows = vec![vec![serde_json::json!(1), serde_json::json!("Ada")]]; @@ -1832,8 +1978,16 @@ mod tests { #[test] fn oracle_import_insert_batches_use_single_row_statements() { let mappings = vec![ - TableImportColumnMapping { source_column: "id".to_string(), target_column: "id".to_string() }, - TableImportColumnMapping { source_column: "name".to_string(), target_column: "name".to_string() }, + TableImportColumnMapping { + source_column: "id".to_string(), + target_column: "id".to_string(), + target_data_type: None, + }, + TableImportColumnMapping { + source_column: "name".to_string(), + target_column: "name".to_string(), + target_data_type: None, + }, ]; let data = ParsedImportFile { columns: vec!["id".to_string(), "name".to_string()], @@ -1873,8 +2027,13 @@ mod tests { TableImportColumnMapping { source_column: "start".to_string(), target_column: "insurance_start_time".to_string(), + target_data_type: None, + }, + TableImportColumnMapping { + source_column: "raw".to_string(), + target_column: "raw_text".to_string(), + target_data_type: None, }, - TableImportColumnMapping { source_column: "raw".to_string(), target_column: "raw_text".to_string() }, ]; let data = ParsedImportFile { columns: vec!["start".to_string(), "raw".to_string()], @@ -1907,8 +2066,11 @@ mod tests { #[test] fn import_insert_batches_preserve_sqlserver_unicode_text() { - let mappings = - vec![TableImportColumnMapping { source_column: "name".to_string(), target_column: "name".to_string() }]; + let mappings = vec![TableImportColumnMapping { + source_column: "name".to_string(), + target_column: "name".to_string(), + target_data_type: None, + }]; let data = ParsedImportFile { columns: vec!["name".to_string()], rows: vec![vec![serde_json::json!("Tiếng Việt")]],