From e515933ab366313112ffac94284f7d116cbec663 Mon Sep 17 00:00:00 2001 From: t8y2 <1156263951@qq.com> Date: Mon, 1 Jun 2026 17:13:42 +0800 Subject: [PATCH] fix(sql): quote postgres completion identifiers --- .../src/components/editor/QueryEditor.vue | 4 + apps/desktop/src/lib/sqlCompletion.ts | 73 +++++++++++++--- packages/app-tests/sqlCompletion.test.ts | 86 +++++++++++++++++++ 3 files changed, 151 insertions(+), 12 deletions(-) diff --git a/apps/desktop/src/components/editor/QueryEditor.vue b/apps/desktop/src/components/editor/QueryEditor.vue index 19ad86483..80e11fac2 100644 --- a/apps/desktop/src/components/editor/QueryEditor.vue +++ b/apps/desktop/src/components/editor/QueryEditor.vue @@ -805,6 +805,7 @@ function buildCompletionResult( label: item.label, type: item.type, detail: item.detail, + apply: item.apply, boost: item.boost, }, ), @@ -835,6 +836,7 @@ async function provideSqlCompletions( schemas: [], translations: completionTranslations.value, snippets: settingsStore.editorSettings.snippets, + dialect: props.dialect, }); return buildCompletionResult(items, position, completionContext.prefix.length, fullDoc); } @@ -853,6 +855,7 @@ async function provideSqlCompletions( schemas: [], translations: completionTranslations.value, snippets: settingsStore.editorSettings.snippets, + dialect: props.dialect, }); return buildCompletionResult(items, position, completionContext.prefix.length, fullDoc); } @@ -1066,6 +1069,7 @@ async function performAsyncCompletionWithResult( schemas: schemaNames, translations: completionTranslations.value, snippets: settingsStore.editorSettings.snippets, + dialect: props.dialect, }); return buildCompletionResult(items, position, completionContext.prefix.length, fullDoc); diff --git a/apps/desktop/src/lib/sqlCompletion.ts b/apps/desktop/src/lib/sqlCompletion.ts index 8096b975f..b4626a4c2 100644 --- a/apps/desktop/src/lib/sqlCompletion.ts +++ b/apps/desktop/src/lib/sqlCompletion.ts @@ -673,6 +673,7 @@ export function buildSqlCompletionItems( columnsByTable: Map; schemas?: string[]; translations?: SqlCompletionTranslations; + dialect?: "mysql" | "postgres" | "sqlserver"; }, ): SqlCompletionItem[] { const context = getSqlCompletionContext(sql, cursor); @@ -687,10 +688,12 @@ export function buildSqlCompletionItemsFromContext( schemas?: string[]; translations?: SqlCompletionTranslations; snippets?: SqlSnippet[]; + dialect?: "mysql" | "postgres" | "sqlserver"; }, ): SqlCompletionItem[] { const items: SqlCompletionItem[] = []; const t = input.translations; + const dialect = input.dialect; if (!context.exclusiveTableSuggestions && !context.exclusiveColumnSuggestions) { items.push(...buildSnippetItems(context.prefix, input.snippets ?? DEFAULT_SQL_SNIPPETS)); @@ -707,11 +710,11 @@ export function buildSqlCompletionItemsFromContext( context.isGroupBy && context.nonAggregatedSelectColumns.length > 0 ) { - items.push(...buildNonAggregatedColumnItems(context, input.columnsByTable)); + items.push(...buildNonAggregatedColumnItems(context, input.columnsByTable, dialect)); } if (!context.exclusiveTableSuggestions && !context.exclusiveColumnSuggestions && context.suggestJoinConditions) { - items.push(...buildJoinConditionItems(context, input.columnsByTable)); + items.push(...buildJoinConditionItems(context, input.columnsByTable, dialect)); } if (context.suggestKeywords) { @@ -719,7 +722,7 @@ export function buildSqlCompletionItemsFromContext( } if (!context.exclusiveTableSuggestions && context.suggestColumns) { - items.push(...buildColumnItems(context, input.columnsByTable)); + items.push(...buildColumnItems(context, input.columnsByTable, dialect)); } // Suggest aliases for referenced tables (independent of table-suggestion mode) @@ -728,9 +731,9 @@ export function buildSqlCompletionItemsFromContext( } if (!context.exclusiveColumnSuggestions && context.suggestTables) { - items.push(...buildTableItems(context.prefix, input.tables)); + items.push(...buildTableItems(context.prefix, input.tables, dialect)); if (input.schemas && input.schemas.length > 0) { - items.push(...buildSchemaItems(context.prefix, input.schemas)); + items.push(...buildSchemaItems(context.prefix, input.schemas, dialect)); } } @@ -741,7 +744,7 @@ export function buildSqlCompletionItemsFromContext( // SELECT * expansion if (context.onStar) { - const starItem = buildStarExpansionItem(input.columnsByTable, t); + const starItem = buildStarExpansionItem(input.columnsByTable, t, dialect); if (starItem) items.push(starItem); } @@ -1485,20 +1488,43 @@ function unquoteIdentifier(value: string): string { return value; } -function buildTableItems(prefix: string, tables: SqlCompletionTable[]): SqlCompletionItem[] { +function quoteSqlIdentifier(identifier: string, dialect?: "mysql" | "postgres" | "sqlserver"): string { + if (dialect !== "postgres" || !requiresPostgresIdentifierQuote(identifier)) return identifier; + return `"${identifier.replaceAll('"', '""')}"`; +} + +function requiresPostgresIdentifierQuote(identifier: string): boolean { + if (!/^[a-z_][a-z0-9_$]*$/.test(identifier)) return true; + return POSTGRES_IDENTIFIER_KEYWORDS.has(identifier); +} + +const POSTGRES_IDENTIFIER_KEYWORDS = new Set( + SQL_KEYWORDS.map((keyword) => keyword.toLowerCase()).concat(["current_user", "session_user", "user"]), +); + +function buildTableItems( + prefix: string, + tables: SqlCompletionTable[], + dialect?: "mysql" | "postgres" | "sqlserver", +): SqlCompletionItem[] { return tables .filter((table) => matchesPrefix(table.name, prefix)) .map((table) => ({ label: table.name, type: "table" as const, detail: table.schema ? `${table.schema}.${table.name}` : table.type, + apply: quoteSqlIdentifier(table.name, dialect), boost: computeBoost(table.name, prefix) + 1000, })) .sort(compareCompletionItems) .slice(0, MAX_TABLE_COMPLETION_ITEMS); } -function buildSchemaItems(prefix: string, schemas: string[]): SqlCompletionItem[] { +function buildSchemaItems( + prefix: string, + schemas: string[], + dialect?: "mysql" | "postgres" | "sqlserver", +): SqlCompletionItem[] { return schemas .filter((schema) => matchesPrefix(schema, prefix)) .slice(0, 50) @@ -1506,7 +1532,7 @@ function buildSchemaItems(prefix: string, schemas: string[]): SqlCompletionItem[ label: schema, type: "schema" as const, detail: "schema", - apply: `${schema}.`, + apply: `${quoteSqlIdentifier(schema, dialect)}.`, boost: computeBoost(schema, prefix) + 1500, })); } @@ -1514,6 +1540,7 @@ function buildSchemaItems(prefix: string, schemas: string[]): SqlCompletionItem[ function buildStarExpansionItem( columnsByTable: Map, t?: SqlCompletionTranslations, + dialect?: "mysql" | "postgres" | "sqlserver", ): SqlCompletionItem | null { const allColumns: string[] = []; const seen = new Set(); @@ -1521,7 +1548,7 @@ function buildStarExpansionItem( for (const col of cols) { if (seen.has(col.name)) continue; seen.add(col.name); - allColumns.push(col.name); + allColumns.push(quoteSqlIdentifier(col.name, dialect)); } } if (allColumns.length === 0) return null; @@ -1698,6 +1725,7 @@ function isInTableListContext(beforeToken: string): boolean { function buildColumnItems( context: SqlCompletionContext, columnsByTable: Map, + dialect?: "mysql" | "postgres" | "sqlserver", ): SqlCompletionItem[] { // Collect all columns from the map (all tables have been fetched) const allColumns: Array = []; @@ -1773,12 +1801,24 @@ function buildColumnItems( label: column.displayLabel, type: "column" as const, detail: buildColumnDetail(column), + apply: buildColumnApply(column, context, dialect), boost: computeBoost(column.displayLabel, context.prefix) + keyBoost, }; }) .sort(compareCompletionItems); } +function buildColumnApply( + column: SqlCompletionColumn & { displayLabel: string }, + context: SqlCompletionContext, + dialect?: "mysql" | "postgres" | "sqlserver", +): string { + if (context.qualifier || !column.displayLabel.includes(".")) { + return quoteSqlIdentifier(column.name, dialect); + } + return `${quoteSqlIdentifier(column.table, dialect)}.${quoteSqlIdentifier(column.name, dialect)}`; +} + function isKeyColumn(name: string): boolean { const lower = name.toLowerCase(); return lower === "id" || lower.endsWith("_id"); @@ -1796,6 +1836,7 @@ function buildColumnDetail(column: SqlCompletionColumn): string { function buildJoinConditionItems( context: SqlCompletionContext, columnsByTable: Map, + dialect?: "mysql" | "postgres" | "sqlserver", ): SqlCompletionItem[] { const refs = context.referencedTables; if (refs.length < 2) return []; @@ -1807,7 +1848,9 @@ function buildJoinConditionItems( for (const previous of previousRefs) { const previousColumns = columnsForReferencedTable(previous, columnsByTable); const latestColumns = columnsForReferencedTable(latest, columnsByTable); - items.push(...buildJoinConditionItemsForPair(previous, previousColumns, latest, latestColumns, context.prefix)); + items.push( + ...buildJoinConditionItemsForPair(previous, previousColumns, latest, latestColumns, context.prefix, dialect), + ); } return items; @@ -1831,10 +1874,13 @@ function buildJoinConditionItemsForPair( right: SqlCompletionReferencedTable, rightColumns: SqlCompletionColumn[], prefix: string, + dialect?: "mysql" | "postgres" | "sqlserver", ): SqlCompletionItem[] { const items: SqlCompletionItem[] = []; const leftRef = left.alias || left.name; const rightRef = right.alias || right.name; + const leftApplyRef = left.alias ? left.alias : quoteSqlIdentifier(left.name, dialect); + const rightApplyRef = right.alias ? right.alias : quoteSqlIdentifier(right.name, dialect); const leftTableKey = singularTableName(left.name); const rightTableKey = singularTableName(right.name); @@ -1884,11 +1930,12 @@ function buildJoinConditionItemsForPair( if (!boost) continue; const label = `${leftLabel} = ${rightLabel}`; if (prefix && !matchesPrefix(label, prefix)) continue; + const apply = `${leftApplyRef}.${quoteSqlIdentifier(leftColumn.name, dialect)} = ${rightApplyRef}.${quoteSqlIdentifier(rightColumn.name, dialect)}`; items.push({ label, type: "snippet", detail: "JOIN condition", - apply: label, + apply, boost, }); } @@ -2004,6 +2051,7 @@ function buildSelectAliasItems(context: SqlCompletionContext): SqlCompletionItem function buildNonAggregatedColumnItems( context: SqlCompletionContext, columnsByTable: Map, + dialect?: "mysql" | "postgres" | "sqlserver", ): SqlCompletionItem[] { const nonAggSet = new Set(context.nonAggregatedSelectColumns.map((c) => c.toLowerCase())); const seen = new Set(); @@ -2019,6 +2067,7 @@ function buildNonAggregatedColumnItems( label: col.name, type: "column" as const, detail: "non-aggregated column — required in GROUP BY", + apply: quoteSqlIdentifier(col.name, dialect), boost: 2800 - items.length, }); } diff --git a/packages/app-tests/sqlCompletion.test.ts b/packages/app-tests/sqlCompletion.test.ts index be605bdee..10706599d 100644 --- a/packages/app-tests/sqlCompletion.test.ts +++ b/packages/app-tests/sqlCompletion.test.ts @@ -38,6 +38,26 @@ const columnsByTable = new Map([ ], ]); +const postgresQuotedTables: SqlCompletionTable[] = [ + { name: "article", schema: "public", type: "table" }, + { name: "order_lines", schema: "public", type: "table" }, + { name: "OrderLines", schema: "public", type: "table" }, + { name: "User", schema: "public", type: "table" }, + { name: 'has"quote', schema: "public", type: "table" }, +]; + +const postgresQuotedColumnsByTable = new Map([ + [ + "public.OrderLines", + [ + { name: "article", table: "OrderLines", schema: "public", dataType: "text" }, + { name: "OrderId", table: "OrderLines", schema: "public", dataType: "uuid" }, + { name: "User", table: "OrderLines", schema: "public", dataType: "text" }, + { name: 'has"quote', table: "OrderLines", schema: "public", dataType: "text" }, + ], + ], +]); + test("suggests SQL keywords for generic keyword input", () => { const items = buildSqlCompletionItems("sel", 3, { tables, @@ -49,6 +69,72 @@ test("suggests SQL keywords for generic keyword input", () => { assert.equal(keyword.type, "keyword"); }); +test("quotes PostgreSQL table identifiers when completion inserts them", () => { + const sql = "select * from Order"; + const items = buildSqlCompletionItems(sql, sql.length, { + tables: postgresQuotedTables, + columnsByTable: new Map(), + dialect: "postgres", + }); + + const table = items.find((item) => item.type === "table" && item.label === "OrderLines"); + assert.equal(table?.apply, '"OrderLines"'); +}); + +test("leaves safe PostgreSQL table identifiers unquoted when completion inserts them", () => { + const sql = "select * from article"; + const items = buildSqlCompletionItems(sql, sql.length, { + tables: postgresQuotedTables, + columnsByTable: new Map(), + dialect: "postgres", + }); + + const table = items.find((item) => item.type === "table" && item.label === "article"); + assert.equal(table?.apply, "article"); +}); + +test("quotes PostgreSQL keyword-like and escaped table identifiers when completion inserts them", () => { + const userItems = buildSqlCompletionItems("select * from User", "select * from User".length, { + tables: postgresQuotedTables, + columnsByTable: new Map(), + dialect: "postgres", + }); + const quotedItems = buildSqlCompletionItems("select * from has", "select * from has".length, { + tables: postgresQuotedTables, + columnsByTable: new Map(), + dialect: "postgres", + }); + + assert.equal(userItems.find((item) => item.label === "User")?.apply, '"User"'); + assert.equal(quotedItems.find((item) => item.label === 'has"quote')?.apply, '"has""quote"'); +}); + +test("quotes PostgreSQL column identifiers when completion inserts them", () => { + const sql = "select Order from public.OrderLines"; + const cursor = "select Order".length; + const items = buildSqlCompletionItems(sql, cursor, { + tables: postgresQuotedTables, + columnsByTable: postgresQuotedColumnsByTable, + dialect: "postgres", + }); + + const column = items.find((item) => item.type === "column" && item.label === "OrderId"); + assert.equal(column?.apply, '"OrderId"'); +}); + +test("leaves safe PostgreSQL column identifiers unquoted when completion inserts them", () => { + const sql = "select arti from public.OrderLines"; + const cursor = "select arti".length; + const items = buildSqlCompletionItems(sql, cursor, { + tables: postgresQuotedTables, + columnsByTable: postgresQuotedColumnsByTable, + dialect: "postgres", + }); + + const column = items.find((item) => item.type === "column" && item.label === "article"); + assert.equal(column?.apply, "article"); +}); + test("suggests matching table names after FROM", () => { const sql = "select * from us"; const items = buildSqlCompletionItems(sql, sql.length, {