fix(sql): quote postgres completion identifiers

This commit is contained in:
t8y2 2026-06-01 17:13:42 +08:00
parent d3b730f900
commit e515933ab3
3 changed files with 151 additions and 12 deletions

View File

@ -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);

View File

@ -673,6 +673,7 @@ export function buildSqlCompletionItems(
columnsByTable: Map<string, SqlCompletionColumn[]>;
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<string, SqlCompletionColumn[]>,
t?: SqlCompletionTranslations,
dialect?: "mysql" | "postgres" | "sqlserver",
): SqlCompletionItem | null {
const allColumns: string[] = [];
const seen = new Set<string>();
@ -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<string, SqlCompletionColumn[]>,
dialect?: "mysql" | "postgres" | "sqlserver",
): SqlCompletionItem[] {
// Collect all columns from the map (all tables have been fetched)
const allColumns: Array<SqlCompletionColumn & { key: string }> = [];
@ -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<string, SqlCompletionColumn[]>,
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<string, SqlCompletionColumn[]>,
dialect?: "mysql" | "postgres" | "sqlserver",
): SqlCompletionItem[] {
const nonAggSet = new Set(context.nonAggregatedSelectColumns.map((c) => c.toLowerCase()));
const seen = new Set<string>();
@ -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,
});
}

View File

@ -38,6 +38,26 @@ const columnsByTable = new Map<string, SqlCompletionColumn[]>([
],
]);
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<string, SqlCompletionColumn[]>([
[
"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, {