fix(sqlserver): honor USE context in completions

This commit is contained in:
guoyongchang 2026-07-31 21:22:31 +08:00 committed by GitHub
parent f64d363abe
commit 486a73ef09
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
16 changed files with 974 additions and 114 deletions

View File

@ -45,7 +45,18 @@ import { buildSqlSemanticModel } from "@/lib/sql/semantic/model";
import { mergeSqlSemanticReferenceAnalysis, resolveSqlSemanticNavigationTarget } from "@/lib/sql/semantic/references";
import { buildElasticsearchCompletionItemsFromContext, getElasticsearchCompletionContext, getElasticsearchCompletionResultValidFor, shouldAutoOpenElasticsearchCompletion, type ElasticsearchCompletionItem } from "@/lib/elasticsearch/elasticsearchCompletion";
import { buildMongoCompletionItemsFromContext, getMongoCompletionContext, getMongoCompletionResultValidFor, mongoCompletionNeedsCollections, mongoCompletionNeedsFields, shouldAutoOpenMongoCompletion, type MongoCompletionItem } from "@/lib/mongo/mongoCompletion";
import { mergeSqlCompletionQualifierNames, resolveSqlCompletionRoutineLookupTarget, resolveSqlCompletionSchemaLookupDatabase, resolveSqlCompletionTableLookupTarget } from "@/lib/sql/sqlCompletionLookupTarget";
import {
buildSqlServerUseDatabaseCompletionItems,
mergeSqlCompletionQualifierNames,
resolveSqlCompletionRoutineLookupTarget,
resolveSqlCompletionSchemaLookupDatabase,
resolveSqlCompletionScope,
resolveSqlCompletionTableLookupTarget,
resolveSqlServerUseDatabaseCompletion,
sqlServerUseCompletionDatabaseNames,
sqlServerUseDatabaseBeforeCursor,
type SqlCompletionScope,
} from "@/lib/sql/sqlCompletionLookupTarget";
import { usesOracleSessionCompletionColumns as shouldUseOracleSessionCompletionColumns } from "@/lib/sql/oracleCompletionSession";
import { extractIdentifierDetailsAt, isSqlKeyword, matchTable, mergeSqlObjectNavigationType, splitQualifiedIdentifier, sqlObjectHoverDetail, sqlObjectNavigationSourceKind, sqlObjectNavigationTarget, type SqlObjectNavigationTarget } from "@/lib/sql/sqlNavigation";
import { buildHoverTableSql, ddlForHoverPreview, hoverTableMatchesScope, quoteQualifiedName, reformatHoverDdl, scopeHoverTables, type HoverTableScope } from "@/lib/editor/hoverTableSql";
@ -102,7 +113,7 @@ import { sqlReferenceAnalysisDialectFor } from "@/lib/sql/semantic/dialect";
import { buildRedisSyntaxDiagnostics, shouldRunRedisDiagnostics } from "@/lib/redis/redisSyntaxDiagnostics";
import { buildRedisCompletionItemsFromContext, getRedisCompletionContext, getRedisCompletionResultValidFor, shouldAutoOpenRedisCompletion, takesKeyArgument, type RedisCompletionItem } from "@/lib/redis/redisCompletion";
import type { SqlCompletionColumn, SqlCompletionForeignKey, SqlCompletionItem, SqlCompletionObject, SqlCompletionReferencedTable, SqlCompletionTable } from "@/lib/sql/sqlCompletion";
import type { CompletionAssistantObjectKind, ColumnInfo, DatabaseType, IndexInfo, SqlReferenceAnalysis, SqlTableReference, SqlTextSpan } from "@/types/database";
import type { CompletionAssistantObjectKind, ColumnInfo, DatabaseType, IndexInfo, SqlReferenceAnalysis, SqlServerCompletionContext, SqlTableReference, SqlTextSpan } from "@/types/database";
const props = defineProps<{
modelValue: string;
@ -422,7 +433,7 @@ function editorThemeAppearance() {
// Completion cache
let cachedTables: SqlCompletionTable[] = [];
let cachedCompletionObjects: SqlCompletionObject[] = [];
const cachedCompletionObjectsByScope = new Map<string, SqlCompletionObject[]>();
// Persistent column cache keyed by "schema.table" or "table"
const cachedColumnsByTable = new Map<string, SqlCompletionColumn[]>();
const cachedInsertValueHintColumnsByTable = new Map<string, string[]>();
@ -1612,9 +1623,12 @@ function identifierRangeAt(sql: string, pos: number): { from: number; to: number
return { from, to, text };
}
function completionCacheKey(table: { name: string; catalog?: string | null; database?: string | null; schema?: string | null }) {
const schema = table.schema ?? props.schema;
const database = supportsDatabaseSchemaQualifierCompletion() ? table.database : undefined;
type CompletionMetadataScope = Pick<SqlCompletionScope, "database" | "schema">;
function completionCacheKey(table: { name: string; catalog?: string | null; database?: string | null; schema?: string | null }, scope?: CompletionMetadataScope) {
const schema = table.schema ?? scope?.schema ?? props.schema;
const scopedDatabase = scope && scope.database !== props.database ? scope.database : undefined;
const database = supportsDatabaseSchemaQualifierCompletion() ? (table.database ?? scopedDatabase) : undefined;
return schema ? `${database ? `${database}.` : ""}${schema}.${table.name}` : table.name;
}
@ -1697,15 +1711,16 @@ function allowsOnDemandQualifiedTableCompletion(prefix: string): boolean {
return prefix.trim().length >= PRESTO_ON_DEMAND_TABLE_COMPLETION_MIN_PREFIX;
}
function completionMetadataTarget(table: { name: string; catalog?: string | null; database?: string | null; schema?: string | null }): { database: string; schema?: string; catalog?: string } | null {
if (props.database == null) return null;
function completionMetadataTarget(table: { name: string; catalog?: string | null; database?: string | null; schema?: string | null }, scope?: CompletionMetadataScope): { database: string; schema?: string; catalog?: string } | null {
const currentDatabase = scope?.database ?? props.database;
if (currentDatabase == null) return null;
if (supportsDatabaseSchemaQualifierCompletion() && table.database) {
return { database: table.database, schema: table.schema ?? undefined, catalog: table.catalog ?? props.catalog };
}
if (supportsDatabaseQualifierCompletion() && table.schema) {
return { database: table.schema, catalog: table.catalog ?? props.catalog };
}
return { database: props.database, schema: table.schema ?? props.schema, catalog: table.catalog ?? props.catalog };
return { database: currentDatabase, schema: table.schema ?? scope?.schema ?? props.schema, catalog: table.catalog ?? props.catalog };
}
function isVirtualCompletionTableReference(table: { name: string; database?: string | null; schema?: string | null }): boolean {
@ -2446,7 +2461,8 @@ function isTypedCompletionActivation(explicit: boolean) {
}
function markCompletionAccepted(item: QueryCompletionItem) {
suppressNextSqlCompletionAutoStartUntil = shouldChainSqlCompletionAfterAccept(item) ? 0 : Date.now() + 750;
const shouldContinueCompletion = shouldChainSqlCompletionAfterAccept(item) || (props.databaseType === "sqlserver" && item.type === "keyword" && item.label.toUpperCase() === "USE");
suppressNextSqlCompletionAutoStartUntil = shouldContinueCompletion ? 0 : Date.now() + 750;
completionEpoch++;
}
@ -2479,15 +2495,17 @@ function mayCompleteDatabaseSchemaQualifier(completionContext: ReturnType<typeof
return (completionContext.qualifierParts?.filter(Boolean).length ?? completionContext.qualifier?.split(".").filter(Boolean).length ?? 0) === 1;
}
function localCompletionSchemasForDatabaseDisambiguation(completionContext: ReturnType<typeof getSqlCompletionContext>, databaseNames: string[]): string[] {
if (!props.connectionId || props.database == null || !mayCompleteDatabaseSchemaQualifier(completionContext)) return [];
function localCompletionSchemasForDatabaseDisambiguation(completionContext: ReturnType<typeof getSqlCompletionContext>, databaseNames: string[], scope?: CompletionMetadataScope): string[] {
const currentDatabase = scope?.database ?? props.database;
const currentSchema = scope?.schema ?? props.schema;
if (!props.connectionId || currentDatabase == null || !mayCompleteDatabaseSchemaQualifier(completionContext)) return [];
const database = resolveSqlCompletionSchemaLookupDatabase({
supportsDatabaseSchemaQualifier: true,
completionContext,
knownDatabases: databaseNames,
});
if (!database) return [];
return mergeSqlCompletionQualifierNames(props.schema ? [props.schema] : [], connectionStore.lookupLocalCompletionSchemas(props.connectionId, props.database, completionContext.qualifier, MAX_COMPLETION_TABLES));
return mergeSqlCompletionQualifierNames(currentSchema ? [currentSchema] : [], connectionStore.lookupLocalCompletionSchemas(props.connectionId, currentDatabase, completionContext.qualifier, MAX_COMPLETION_TABLES));
}
function shouldInsertSqlCompletionSpace(): boolean {
@ -2672,11 +2690,44 @@ async function provideSqlCompletions(context: CompletionContext) {
try {
if (isSqlCompletionSuppressedContext(fullDoc, position)) return null;
if (!explicit && !shouldAutoOpenSqlCompletion(fullDoc, position, sqlCompletionDialectOptions())) return null;
const useDatabaseCompletion = resolveSqlServerUseDatabaseCompletion({
sql: fullDoc,
cursor: position,
databaseType: props.databaseType,
});
if (!explicit && !useDatabaseCompletion && !shouldAutoOpenSqlCompletion(fullDoc, position, sqlCompletionDialectOptions())) return null;
if (useDatabaseCompletion) {
const currentDatabase = props.database ?? "";
if (!currentDatabase) return null;
let sqlServerContext: SqlServerCompletionContext;
try {
sqlServerContext = await connectionStore.getSqlServerCompletionContext(props.connectionId, currentDatabase);
} catch {
// Without a server-reported capability, do not suggest a USE target
// that the current SQL Server session may be unable to switch to.
return null;
}
if (sqlServerContext.supports_session_database_switch) {
try {
await connectionStore.listCompletionDatabases(props.connectionId);
} catch {
// Keep locally indexed database names available when metadata refresh fails.
}
}
if (epoch !== completionEpoch) return null;
const databaseNames = sqlServerUseCompletionDatabaseNames({
databaseNames: connectionStore.lookupLocalCompletionDatabases(props.connectionId, useDatabaseCompletion.prefix, MAX_COMPLETION_TABLES),
currentDatabase,
supportsSessionDatabaseSwitch: sqlServerContext.supports_session_database_switch,
});
const items = buildSqlServerUseDatabaseCompletionItems(databaseNames, useDatabaseCompletion);
return buildCompletionResult(items, useDatabaseCompletion.from);
}
const legacyCompletionContext = getSqlCompletionContext(fullDoc, position, sqlCompletionDialectOptions());
const semanticModel = SEMANTIC_SQL_COMPLETION_ENABLED ? buildSqlSemanticModel(fullDoc, position, sqlCompletionDialectOptions()) : null;
const completionContext = semanticModel ? sqlCompletionContextFromSemantic(semanticModel, legacyCompletionContext) : legacyCompletionContext;
let completionContext = semanticModel ? sqlCompletionContextFromSemantic(semanticModel, legacyCompletionContext) : legacyCompletionContext;
if (!hasDatabase) {
const items = buildSqlCompletionItemsFromContext(completionContext, {
@ -2695,6 +2746,45 @@ async function provideSqlCompletions(context: CompletionContext) {
return buildCompletionResult(items, position - completionContext.prefix.length, getSqlCompletionResultValidFor(fullDoc, position));
}
const useDatabase = props.databaseType === "sqlserver" ? sqlServerUseDatabaseBeforeCursor(fullDoc, position) : undefined;
let knownUseDatabases: string[] | undefined;
let supportsSessionDatabaseSwitch: boolean | undefined;
let useDatabaseDefaultSchema: string | undefined;
if (useDatabase) {
try {
const currentContext = await connectionStore.getSqlServerCompletionContext(props.connectionId, props.database!);
supportsSessionDatabaseSwitch = currentContext.supports_session_database_switch;
knownUseDatabases = [props.database!];
if (supportsSessionDatabaseSwitch) {
knownUseDatabases = mergeSqlCompletionQualifierNames(knownUseDatabases, connectionStore.lookupLocalCompletionDatabases(props.connectionId, "", MAX_COMPLETION_TABLES));
if (!knownUseDatabases.some((database) => database.toLowerCase() === useDatabase.toLowerCase())) {
knownUseDatabases = mergeSqlCompletionQualifierNames(knownUseDatabases, await connectionStore.listCompletionDatabases(props.connectionId));
}
}
const targetDatabase = knownUseDatabases.find((database) => database.toLowerCase() === useDatabase.toLowerCase());
if (targetDatabase) {
const targetContext = targetDatabase.toLowerCase() === props.database!.toLowerCase() ? currentContext : await connectionStore.getSqlServerCompletionContext(props.connectionId, targetDatabase);
useDatabaseDefaultSchema = targetContext.default_schema;
}
} catch {
// An unverified USE target must not replace the selected database.
}
if (epoch !== completionEpoch) return null;
}
const completionScope = resolveSqlCompletionScope({
sql: fullDoc,
cursor: position,
databaseType: props.databaseType,
currentDatabase: props.database!,
currentSchema: props.schema,
knownDatabases: knownUseDatabases,
supportsSessionDatabaseSwitch,
useDatabaseDefaultSchema,
completionContext,
});
completionContext = completionScope.completionContext;
const needsAsyncData =
completionContext.suggestTables || completionContext.suggestRoutines || completionContext.exclusiveRoutineSuggestions || !!completionContext.qualifier || !!completionContext.insertTable || completionContext.exclusiveColumnSuggestions || completionContext.referencedTables.length > 0;
@ -2718,14 +2808,14 @@ async function provideSqlCompletions(context: CompletionContext) {
const tableNameCompletion = isTableNameCompletionContext(completionContext);
const shouldResolveColumnCompletion = completionContext.suggestColumns && completionContext.referencedTables.length > 0 && (completionContext.prefix.length > 0 || typedActivation);
const shouldResolveAsyncCompletion = tableNameCompletion || shouldResolveColumnCompletion;
const localResult = buildLocalSqlCompletionResult(completionContext, fullDoc, position);
const localResult = buildLocalSqlCompletionResult(completionContext, fullDoc, position, completionScope);
if (localResult) {
scheduleCompletionMetadataRefresh(completionContext, fullDoc, position);
scheduleCompletionMetadataRefresh(completionContext, fullDoc, position, completionScope);
const hasLocalColumnResult = localResult.options.some((option) => option.type === "column");
if ((!explicit || typedActivation) && (!shouldResolveColumnCompletion || hasLocalColumnResult)) return localResult;
}
if ((!explicit || typedActivation) && !shouldResolveAsyncCompletion) {
scheduleCompletionMetadataRefresh(completionContext, fullDoc, position);
scheduleCompletionMetadataRefresh(completionContext, fullDoc, position, completionScope);
return null;
}
@ -2749,7 +2839,7 @@ async function provideSqlCompletions(context: CompletionContext) {
return;
}
try {
const result = await performAsyncCompletionWithResult(epoch, completionContext, fullDoc, position);
const result = await performAsyncCompletionWithResult(epoch, completionContext, fullDoc, position, completionScope);
resolve(result ?? localResult);
} catch {
resolve(localResult);
@ -2785,7 +2875,9 @@ function flushImeComposition() {
emit("cursorChange", currentView.state.selection.main.head);
latestSelection = readEditorSelection(currentView);
if (editorIsActive) emitEditorSelection(latestSelection);
if (shouldAutoOpenSqlCompletion(currentView.state.doc.toString(), currentView.state.selection.main.head, sqlCompletionDialectOptions())) {
const fullDoc = currentView.state.doc.toString();
const position = currentView.state.selection.main.head;
if (resolveSqlServerUseDatabaseCompletion({ sql: fullDoc, cursor: position, databaseType: props.databaseType }) || shouldAutoOpenSqlCompletion(fullDoc, position, sqlCompletionDialectOptions())) {
scheduleSqlCompletionStart(currentView);
}
}
@ -2793,6 +2885,7 @@ function flushImeComposition() {
function shouldStartSqlCompletionAfterInput(insertedText: string, removedText: string, currentView: EditorViewType): boolean {
const position = currentView.state.selection.main.head;
const fullDoc = currentView.state.doc.toString();
if (resolveSqlServerUseDatabaseCompletion({ sql: fullDoc, cursor: position, databaseType: props.databaseType })) return true;
if (!insertedText && removedText) {
const completionContext = getSqlCompletionContext(fullDoc, position, sqlCompletionDialectOptions());
return isTableNameCompletionContext(completionContext) && shouldAutoOpenSqlCompletion(fullDoc, position, sqlCompletionDialectOptions());
@ -2810,10 +2903,10 @@ function shouldStartSqlCompletionAfterInput(insertedText: string, removedText: s
return isTableNameCompletionContext(completionContext) || shouldAutoOpenSqlCompletion(fullDoc, position, sqlCompletionDialectOptions());
}
function buildLocalSqlCompletionResult(completionContext: ReturnType<typeof getSqlCompletionContext>, fullDoc: string, position: number) {
function buildLocalSqlCompletionResult(completionContext: ReturnType<typeof getSqlCompletionContext>, fullDoc: string, position: number, scope: CompletionMetadataScope) {
if (!props.connectionId || props.database == null) return null;
const databaseNames = localCompletionDatabaseNames(completionContext);
const currentDatabaseSchemaNames = localCompletionSchemasForDatabaseDisambiguation(completionContext, databaseNames);
const currentDatabaseSchemaNames = localCompletionSchemasForDatabaseDisambiguation(completionContext, databaseNames, scope);
const schemaLookupDatabase = resolveSqlCompletionSchemaLookupDatabase({
supportsDatabaseSchemaQualifier: supportsDatabaseSchemaQualifierCompletion(),
completionContext,
@ -2822,8 +2915,8 @@ function buildLocalSqlCompletionResult(completionContext: ReturnType<typeof getS
});
const shouldLoadTables = !schemaLookupDatabase && (completionContext.suggestTables || (!!completionContext.qualifier && !isReferencedTableQualifier(completionContext)));
const tableLookupTarget = resolveSqlCompletionTableLookupTarget({
currentDatabase: props.database,
currentSchema: props.schema,
currentDatabase: scope.database,
currentSchema: scope.schema,
supportsDatabaseQualifier: supportsDatabaseQualifierCompletion(),
supportsDatabaseSchemaQualifier: supportsDatabaseSchemaQualifierCompletion(),
completionContext,
@ -2833,35 +2926,37 @@ function buildLocalSqlCompletionResult(completionContext: ReturnType<typeof getS
const tables = schemaLookupDatabase ? [] : shouldLoadTables ? connectionStore.lookupLocalCompletionTables(props.connectionId, tableLookupTarget.database, tableLookupTarget.filter, MAX_COMPLETION_TABLES, globalOracleTableSearch ? undefined : tableLookupTarget.schema, props.catalog) : cachedTables;
const shouldLoadObjects = shouldLoadCompletionObjects(completionContext);
const completionObjects = shouldLoadObjects ? lookupLocalCompletionObjectsForContext(completionContext) : cachedCompletionObjects;
const completionObjectScope = routineCompletionScopeForContext(completionContext, scope);
const scopedCachedCompletionObjects = completionObjectsForScope(completionObjectScope);
const completionObjects = shouldLoadObjects ? lookupLocalCompletionObjectsForContext(completionContext, scope) : scopedCachedCompletionObjects;
const schemaNames =
completionContext.suggestTables && !completionContext.insertTable
? schemaLookupDatabase
? connectionStore.lookupLocalCompletionSchemas(props.connectionId, schemaLookupDatabase, completionContext.prefix, MAX_COMPLETION_TABLES)
: !completionContext.qualifier
? mergeSqlCompletionQualifierNames(connectionStore.lookupLocalCompletionSchemas(props.connectionId, props.database, completionContext.prefix, MAX_COMPLETION_TABLES), databaseNames)
? mergeSqlCompletionQualifierNames(connectionStore.lookupLocalCompletionSchemas(props.connectionId, scope.database, completionContext.prefix, MAX_COMPLETION_TABLES), databaseNames)
: []
: [];
const columnsByTable = new Map<string, SqlCompletionColumn[]>();
if (completionContext.insertTable) {
const insertDatabase = (supportsDatabaseSchemaQualifierCompletion() ? completionContext.insertDatabase : undefined) ?? props.database;
const insertSchema = completionContext.insertSchema ?? props.schema;
const insertDatabase = (supportsDatabaseSchemaQualifierCompletion() ? completionContext.insertDatabase : undefined) ?? scope.database;
const insertSchema = completionContext.insertSchema ?? scope.schema;
const insertColumns = usesOracleSessionCompletionColumns(insertSchema) ? [] : connectionStore.lookupLocalCompletionColumns(props.connectionId, insertDatabase, completionContext.insertTable, insertSchema, props.catalog);
if (insertColumns.length > 0) {
columnsByTable.set(completionCacheKey({ name: completionContext.insertTable, database: completionContext.insertDatabase, schema: insertSchema }), insertColumns);
columnsByTable.set(completionCacheKey({ name: completionContext.insertTable, database: completionContext.insertDatabase, schema: insertSchema }, scope), insertColumns);
}
}
const qualifiedColumnTarget = completionQualifiedTableTarget(completionContext);
if (qualifiedColumnTarget) {
const cacheKey = completionCacheKey(qualifiedColumnTarget);
const cacheKey = completionCacheKey(qualifiedColumnTarget, scope);
const cached = cachedColumnsByTable.get(cacheKey);
if (cached) {
columnsByTable.set(cacheKey, cached);
} else {
const target = completionMetadataTarget(qualifiedColumnTarget);
const target = completionMetadataTarget(qualifiedColumnTarget, scope);
const localColumns = target && !usesOracleSessionCompletionColumns(target.schema) ? connectionStore.lookupLocalCompletionColumns(props.connectionId, target.database, qualifiedColumnTarget.name, target.schema, target.catalog) : [];
if (localColumns.length > 0) {
columnsByTable.set(cacheKey, localColumns);
@ -2884,13 +2979,13 @@ function buildLocalSqlCompletionResult(completionContext: ReturnType<typeof getS
);
continue;
}
const cacheKey = completionCacheKey(refTable);
const cacheKey = completionCacheKey(refTable, scope);
const cached = cachedColumnsByTable.get(cacheKey);
if (cached) {
columnsByTable.set(cacheKey, cached);
continue;
}
const target = completionMetadataTarget(refTable);
const target = completionMetadataTarget(refTable, scope);
const localColumns = target && !usesOracleSessionCompletionColumns(target.schema) ? connectionStore.lookupLocalCompletionColumns(props.connectionId, target.database, refTable.name, target.schema, target.catalog, refTable) : [];
if (localColumns.length > 0) {
columnsByTable.set(cacheKey, localColumns);
@ -2915,7 +3010,7 @@ function buildLocalSqlCompletionResult(completionContext: ReturnType<typeof getS
snippets: settingsStore.editorSettings.snippets,
dialect: props.dialect,
databaseType: snippetDatabaseType.value,
currentSchema: props.schema,
currentSchema: scope.schema,
keywordCase: settingsStore.editorSettings.sqlFormatter.keywordCase,
autoAliasTables: settingsStore.editorSettings.autoAliasTables,
});
@ -2923,15 +3018,15 @@ function buildLocalSqlCompletionResult(completionContext: ReturnType<typeof getS
return buildCompletionResult(items, position - completionContext.prefix.length, getSqlCompletionResultValidFor(fullDoc, position));
}
function scheduleCompletionMetadataRefresh(completionContext: ReturnType<typeof getSqlCompletionContext>, fullDoc: string, position: number) {
function scheduleCompletionMetadataRefresh(completionContext: ReturnType<typeof getSqlCompletionContext>, fullDoc: string, position: number, scope: CompletionMetadataScope) {
if (!props.connectionId || props.database == null) return;
const localOnlyMetadata = usesLocalOnlyCompletionMetadata();
const onDemandOnlyColumns = usesOnDemandOnlyCompletionColumns();
const tableNameCompletion = isTableNameCompletionContext(completionContext);
const connectionId = props.connectionId;
const database = props.database;
const database = scope.database;
const databaseNames = localCompletionDatabaseNames(completionContext);
const currentDatabaseSchemaNames = localCompletionSchemasForDatabaseDisambiguation(completionContext, databaseNames);
const currentDatabaseSchemaNames = localCompletionSchemasForDatabaseDisambiguation(completionContext, databaseNames, scope);
const schemaLookupDatabase = resolveSqlCompletionSchemaLookupDatabase({
supportsDatabaseSchemaQualifier: supportsDatabaseSchemaQualifierCompletion(),
completionContext,
@ -2940,7 +3035,7 @@ function scheduleCompletionMetadataRefresh(completionContext: ReturnType<typeof
});
const tableLookupTarget = resolveSqlCompletionTableLookupTarget({
currentDatabase: database,
currentSchema: props.schema,
currentSchema: scope.schema,
supportsDatabaseQualifier: supportsDatabaseQualifierCompletion(),
supportsDatabaseSchemaQualifier: supportsDatabaseSchemaQualifierCompletion(),
completionContext,
@ -2949,7 +3044,7 @@ function scheduleCompletionMetadataRefresh(completionContext: ReturnType<typeof
if (!localOnlyMetadata && !schemaLookupDatabase && (completionContext.suggestTables || (!!completionContext.qualifier && !isReferencedTableQualifier(completionContext)))) {
const globalOracleTableSearch = props.databaseType === "oracle" && completionContext.suggestTables && !completionContext.qualifier;
void connectionStore
.refreshCompletionTables(connectionId, tableLookupTarget.database, tableLookupTarget.filter, MAX_COMPLETION_TABLES, tableLookupTarget.schema, globalOracleTableSearch, props.schema, props.catalog)
.refreshCompletionTables(connectionId, tableLookupTarget.database, tableLookupTarget.filter, MAX_COMPLETION_TABLES, tableLookupTarget.schema, globalOracleTableSearch, scope.schema, props.catalog)
.then((tables) => {
const scopedTables = tables.map((table) => ({ ...table, database: table.database ?? tableLookupTarget.database }));
cachedTables = mergeCompletionTables(cachedTables, scopedTables);
@ -2960,11 +3055,13 @@ function scheduleCompletionMetadataRefresh(completionContext: ReturnType<typeof
.catch(() => {});
}
if (!localOnlyMetadata && shouldLoadCompletionObjects(completionContext)) {
void listCompletionObjectsForContext(completionContext)
const completionObjectScope = routineCompletionScopeForContext(completionContext, scope);
void listCompletionObjectsForContext(completionContext, scope)
.then((objects) => {
const merged = mergeCompletionObjects(cachedCompletionObjects, objects);
const changed = completionObjectsDiffer(cachedCompletionObjects, merged);
cachedCompletionObjects = merged;
const cachedObjects = completionObjectsForScope(completionObjectScope);
const merged = mergeCompletionObjects(cachedObjects, objects);
const changed = completionObjectsDiffer(cachedObjects, merged);
cachedCompletionObjectsByScope.set(completionObjectScopeKey(completionObjectScope), merged);
if (changed) refreshActiveSqlCompletion(fullDoc, position, completionContext);
})
.catch(() => {});
@ -2982,17 +3079,17 @@ function scheduleCompletionMetadataRefresh(completionContext: ReturnType<typeof
if (!onDemandOnlyColumns && completionContext.insertTable) {
const insertTable = completionContext.insertTable;
const insertDatabase = (supportsDatabaseSchemaQualifierCompletion() ? completionContext.insertDatabase : undefined) ?? database;
void refreshCompletionColumnsForEditor(connectionId, insertDatabase, insertTable, completionContext.insertSchema ?? props.schema)
void refreshCompletionColumnsForEditor(connectionId, insertDatabase, insertTable, completionContext.insertSchema ?? scope.schema)
.then((columns) => {
const insertSchema = completionContext.insertSchema ?? props.schema;
cachedColumnsByTable.set(completionCacheKey({ name: insertTable, database: completionContext.insertDatabase, schema: insertSchema }), columns);
const insertSchema = completionContext.insertSchema ?? scope.schema;
cachedColumnsByTable.set(completionCacheKey({ name: insertTable, database: completionContext.insertDatabase, schema: insertSchema }, scope), columns);
})
.catch(() => {});
}
const qualifiedColumnTarget = completionQualifiedTableTarget(completionContext);
const qualifiedColumnCacheKey = qualifiedColumnTarget ? completionCacheKey(qualifiedColumnTarget) : undefined;
const qualifiedColumnCacheKey = qualifiedColumnTarget ? completionCacheKey(qualifiedColumnTarget, scope) : undefined;
if (!onDemandOnlyColumns && qualifiedColumnTarget && qualifiedColumnCacheKey && !cachedColumnsByTable.has(qualifiedColumnCacheKey)) {
const target = completionMetadataTarget(qualifiedColumnTarget);
const target = completionMetadataTarget(qualifiedColumnTarget, scope);
if (target) {
void refreshCompletionColumnsForEditor(connectionId, target.database, qualifiedColumnTarget.name, target.schema, target.catalog)
.then((columns) => {
@ -3005,10 +3102,10 @@ function scheduleCompletionMetadataRefresh(completionContext: ReturnType<typeof
for (const refTable of completionContext.referencedTables) {
if (isVirtualCompletionTableReference(refTable)) continue;
if (refTable.columns && refTable.columns.length > 0) continue;
const cacheKey = completionCacheKey(refTable);
const cacheKey = completionCacheKey(refTable, scope);
if (cacheKey === qualifiedColumnCacheKey) continue;
if (cachedColumnsByTable.has(cacheKey)) continue;
const target = completionMetadataTarget(refTable);
const target = completionMetadataTarget(refTable, scope);
if (!target) continue;
void refreshCompletionColumnsForEditor(connectionId, target.database, refTable.name, target.schema, target.catalog, refTable)
.then((columns) => {
@ -3048,9 +3145,9 @@ function withCompletionLatencyBudget<T>(remote: Promise<T>, local: T): Promise<T
return Promise.race([remote, new Promise<T>((resolve) => setTimeout(() => resolve(local), COMPLETION_REMOTE_LATENCY_BUDGET_MS))]);
}
function listCompletionTablesWithLatencyBudget(connectionId: string, database: string, filter: string, limit: number, schema?: string, globalSearch = false, catalog = props.catalog): Promise<SqlCompletionTable[]> {
function listCompletionTablesWithLatencyBudget(connectionId: string, database: string, filter: string, limit: number, schema?: string, globalSearch = false, catalog = props.catalog, currentSchema = props.schema): Promise<SqlCompletionTable[]> {
const local = connectionStore.lookupLocalCompletionTables(connectionId, database, filter, limit, globalSearch ? undefined : schema, catalog).map((table) => ({ ...table, catalog: table.catalog ?? catalog, database: table.database ?? database }));
const remote = connectionStore.listCompletionTables(connectionId, database, filter, limit, schema, globalSearch, props.schema, catalog).then((tables) => {
const remote = connectionStore.listCompletionTables(connectionId, database, filter, limit, schema, globalSearch, currentSchema, catalog).then((tables) => {
const scopedTables = tables.map((table) => ({ ...table, catalog: table.catalog ?? catalog, database: table.database ?? database }));
cachedTables = mergeCompletionTables(cachedTables, scopedTables);
return scopedTables;
@ -3083,24 +3180,39 @@ function oracleRoutineCompletionTargets(completionContext: ReturnType<typeof get
return [{ schema: parts[parts.length - 2], parentName: parts[parts.length - 1] }];
}
function lookupLocalCompletionObjectsForContext(completionContext: ReturnType<typeof getSqlCompletionContext>): SqlCompletionObject[] {
if (!props.connectionId || props.database == null) return [];
if (props.databaseType === "oracle") {
return connectionStore.lookupLocalCompletionObjects(props.connectionId, props.database, completionContext.prefix, MAX_COMPLETION_TABLES);
}
const target = resolveSqlCompletionRoutineLookupTarget({ currentSchema: props.schema, completionContext });
return connectionStore.lookupLocalCompletionObjects(props.connectionId, props.database, target.mask, MAX_COMPLETION_TABLES, target.schema);
function routineCompletionTargetForContext(completionContext: ReturnType<typeof getSqlCompletionContext>, scope: CompletionMetadataScope) {
return resolveSqlCompletionRoutineLookupTarget({
currentDatabase: scope.database,
currentSchema: scope.schema,
supportsDatabaseSchemaQualifier: supportsDatabaseSchemaQualifierCompletion(),
completionContext,
});
}
async function listCompletionObjectsForContext(completionContext: ReturnType<typeof getSqlCompletionContext>): Promise<SqlCompletionObject[]> {
function routineCompletionScopeForContext(completionContext: ReturnType<typeof getSqlCompletionContext>, scope: CompletionMetadataScope): CompletionMetadataScope {
if (props.databaseType === "oracle") return scope;
const target = routineCompletionTargetForContext(completionContext, scope);
return { database: target.database, schema: target.schema };
}
function lookupLocalCompletionObjectsForContext(completionContext: ReturnType<typeof getSqlCompletionContext>, scope: CompletionMetadataScope): SqlCompletionObject[] {
if (!props.connectionId || props.database == null) return [];
if (props.databaseType === "oracle") {
return connectionStore.lookupLocalCompletionObjects(props.connectionId, scope.database, completionContext.prefix, MAX_COMPLETION_TABLES);
}
const target = routineCompletionTargetForContext(completionContext, scope);
return connectionStore.lookupLocalCompletionObjects(props.connectionId, target.database, target.mask, MAX_COMPLETION_TABLES, target.schema);
}
async function listCompletionObjectsForContext(completionContext: ReturnType<typeof getSqlCompletionContext>, scope: CompletionMetadataScope): Promise<SqlCompletionObject[]> {
if (!props.connectionId || props.database == null) return [];
const objectKinds = completionObjectKindsForContext(completionContext);
if (props.databaseType !== "oracle") {
const target = resolveSqlCompletionRoutineLookupTarget({ currentSchema: props.schema, completionContext });
return connectionStore.listCompletionObjects(props.connectionId, props.database, target.mask, MAX_COMPLETION_TABLES, target.schema, undefined, false, props.schema, objectKinds);
const target = routineCompletionTargetForContext(completionContext, scope);
return connectionStore.listCompletionObjects(props.connectionId, target.database, target.mask, MAX_COMPLETION_TABLES, target.schema, undefined, false, scope.schema, objectKinds);
}
const groups = await Promise.all(
oracleRoutineCompletionTargets(completionContext).map((target) => connectionStore.listCompletionObjects(props.connectionId!, props.database!, completionContext.prefix, MAX_COMPLETION_TABLES, target.schema, target.parentName, target.globalSearch, props.schema, objectKinds)),
oracleRoutineCompletionTargets(completionContext).map((target) => connectionStore.listCompletionObjects(props.connectionId!, scope.database, completionContext.prefix, MAX_COMPLETION_TABLES, target.schema, target.parentName, target.globalSearch, scope.schema, objectKinds)),
);
return groups.reduce((objects, group) => mergeCompletionObjects(objects, group), [] as SqlCompletionObject[]);
}
@ -3111,19 +3223,19 @@ function completionObjectKindsForContext(completionContext: ReturnType<typeof ge
return ["routine"];
}
async function performAsyncCompletionWithResult(epoch: number, completionContext: ReturnType<typeof getSqlCompletionContext>, fullDoc: string, position: number) {
async function performAsyncCompletionWithResult(epoch: number, completionContext: ReturnType<typeof getSqlCompletionContext>, fullDoc: string, position: number, scope: CompletionMetadataScope) {
const localOnlyMetadata = usesLocalOnlyCompletionMetadata();
const onDemandOnlyColumns = usesOnDemandOnlyCompletionColumns();
// Handle INSERT column list: fetch columns for the target table
let insertColumnsByTable = new Map<string, SqlCompletionColumn[]>();
if (completionContext.insertTable) {
try {
const insertDatabase = (supportsDatabaseSchemaQualifierCompletion() ? completionContext.insertDatabase : undefined) ?? props.database!;
const insertCols = await listCompletionColumnsForEditor(props.connectionId!, insertDatabase, completionContext.insertTable, completionContext.insertSchema ?? props.schema);
const insertDatabase = (supportsDatabaseSchemaQualifierCompletion() ? completionContext.insertDatabase : undefined) ?? scope.database;
const insertCols = await listCompletionColumnsForEditor(props.connectionId!, insertDatabase, completionContext.insertTable, completionContext.insertSchema ?? scope.schema);
if (epoch !== completionEpoch) return null;
if (insertCols.length > 0) {
const insertSchema = completionContext.insertSchema ?? props.schema;
const insertKey = completionCacheKey({ name: completionContext.insertTable, database: completionContext.insertDatabase, schema: insertSchema });
const insertSchema = completionContext.insertSchema ?? scope.schema;
const insertKey = completionCacheKey({ name: completionContext.insertTable, database: completionContext.insertDatabase, schema: insertSchema }, scope);
insertColumnsByTable.set(insertKey, insertCols);
}
} catch {
@ -3132,12 +3244,12 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
}
let databaseNames = localCompletionDatabaseNames(completionContext);
let currentDatabaseSchemaNames = localCompletionSchemasForDatabaseDisambiguation(completionContext, databaseNames);
let currentDatabaseSchemaNames = localCompletionSchemasForDatabaseDisambiguation(completionContext, databaseNames, scope);
const mayCompleteDatabaseSchema = mayCompleteDatabaseSchemaQualifier(completionContext);
if (!localOnlyMetadata && supportsDatabaseNameCompletion(props.databaseType) && completionContext.suggestTables && !completionContext.insertTable && (!completionContext.qualifier || mayCompleteDatabaseSchema)) {
const [databasesResult, schemasResult] = await Promise.allSettled([connectionStore.listCompletionDatabases(props.connectionId!), mayCompleteDatabaseSchema ? connectionStore.listCompletionSchemas(props.connectionId!, props.database!) : Promise.resolve(currentDatabaseSchemaNames)]);
const [databasesResult, schemasResult] = await Promise.allSettled([connectionStore.listCompletionDatabases(props.connectionId!), mayCompleteDatabaseSchema ? connectionStore.listCompletionSchemas(props.connectionId!, scope.database) : Promise.resolve(currentDatabaseSchemaNames)]);
databaseNames = databasesResult.status === "fulfilled" ? databasesResult.value : [];
if (schemasResult.status === "fulfilled") currentDatabaseSchemaNames = mergeSqlCompletionQualifierNames(props.schema ? [props.schema] : [], schemasResult.value);
if (schemasResult.status === "fulfilled") currentDatabaseSchemaNames = mergeSqlCompletionQualifierNames(scope.schema ? [scope.schema] : [], schemasResult.value);
if (epoch !== completionEpoch) return null;
}
const schemaLookupDatabase = resolveSqlCompletionSchemaLookupDatabase({
@ -3148,8 +3260,8 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
});
const shouldLoadTables = !schemaLookupDatabase && (completionContext.suggestTables || (!!completionContext.qualifier && !isReferencedTableQualifier(completionContext)));
const tableLookupTarget = resolveSqlCompletionTableLookupTarget({
currentDatabase: props.database!,
currentSchema: props.schema,
currentDatabase: scope.database,
currentSchema: scope.schema,
supportsDatabaseQualifier: supportsDatabaseQualifierCompletion(),
supportsDatabaseSchemaQualifier: supportsDatabaseSchemaQualifierCompletion(),
completionContext,
@ -3161,30 +3273,33 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
: shouldLoadTables
? localOnlyMetadata
? connectionStore.lookupLocalCompletionTables(props.connectionId!, tableLookupTarget.database, tableLookupTarget.filter, MAX_COMPLETION_TABLES, globalOracleTableSearch ? undefined : tableLookupTarget.schema, props.catalog)
: await listCompletionTablesWithLatencyBudget(props.connectionId!, tableLookupTarget.database, tableLookupTarget.filter, MAX_COMPLETION_TABLES, tableLookupTarget.schema, globalOracleTableSearch)
: await listCompletionTablesWithLatencyBudget(props.connectionId!, tableLookupTarget.database, tableLookupTarget.filter, MAX_COMPLETION_TABLES, tableLookupTarget.schema, globalOracleTableSearch, props.catalog, scope.schema)
: cachedTables;
if (localOnlyMetadata && tables.length === 0 && supportsDatabaseSchemaQualifierCompletion() && (completionContext.qualifierParts?.length ?? 0) >= 2 && allowsOnDemandQualifiedTableCompletion(completionContext.prefix)) {
tables = await listCompletionTablesWithLatencyBudget(props.connectionId!, tableLookupTarget.database, tableLookupTarget.filter, PRESTO_ON_DEMAND_TABLE_COMPLETION_LIMIT, tableLookupTarget.schema);
tables = await listCompletionTablesWithLatencyBudget(props.connectionId!, tableLookupTarget.database, tableLookupTarget.filter, PRESTO_ON_DEMAND_TABLE_COMPLETION_LIMIT, tableLookupTarget.schema, false, props.catalog, scope.schema);
}
if (epoch !== completionEpoch) return null;
const shouldLoadObjects = shouldLoadCompletionObjects(completionContext);
let completionObjects = shouldLoadObjects ? (localOnlyMetadata ? lookupLocalCompletionObjectsForContext(completionContext) : await listCompletionObjectsForContext(completionContext)) : cachedCompletionObjects;
const completionObjectScope = routineCompletionScopeForContext(completionContext, scope);
const scopedCachedCompletionObjects = completionObjectsForScope(completionObjectScope);
let completionObjects = shouldLoadObjects ? (localOnlyMetadata ? lookupLocalCompletionObjectsForContext(completionContext, scope) : await listCompletionObjectsForContext(completionContext, scope)) : scopedCachedCompletionObjects;
if (epoch !== completionEpoch) return null;
if (!props.catalog && props.databaseType !== "oracle" && !localOnlyMetadata && completionContext.qualifier && completionObjects.length === 0) {
const schemaObjects = await connectionStore.listCompletionObjects(props.connectionId!, props.database!, completionContext.prefix, MAX_COMPLETION_TABLES, completionContext.qualifier);
const target = routineCompletionTargetForContext(completionContext, scope);
const schemaObjects = await connectionStore.listCompletionObjects(props.connectionId!, target.database, target.mask, MAX_COMPLETION_TABLES, target.schema, undefined, false, scope.schema);
if (schemaObjects.length > 0) {
completionObjects = schemaObjects;
}
if (epoch !== completionEpoch) return null;
}
cachedCompletionObjects = mergeCompletionObjects(cachedCompletionObjects, completionObjects);
cachedCompletionObjectsByScope.set(completionObjectScopeKey(completionObjectScope), mergeCompletionObjects(scopedCachedCompletionObjects, completionObjects));
// Fetch schemas for schema completion
let schemaNames: string[] = [];
if (completionContext.suggestTables && !completionContext.insertTable && (schemaLookupDatabase || !completionContext.qualifier)) {
const database = schemaLookupDatabase ?? props.database!;
const database = schemaLookupDatabase ?? scope.database;
if (localOnlyMetadata) {
const schemas = connectionStore.lookupLocalCompletionSchemas(props.connectionId!, database, completionContext.prefix, MAX_COMPLETION_TABLES);
schemaNames = schemaLookupDatabase ? schemas : mergeSqlCompletionQualifierNames(schemas, databaseNames);
@ -3202,11 +3317,11 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
// If qualifier didn't match any table names, try it as a schema name
let qualifierIsSchema = false;
if (completionContext.qualifier && !schemaLookupDatabase && !tableLookupTarget.qualifierDatabase && !isReferencedTableQualifier(completionContext) && tables.length === 0 && (completionContext.suggestTables || completionContext.exclusiveColumnSuggestions)) {
let schemaTables = connectionStore.lookupLocalCompletionTables(props.connectionId!, props.database!, completionContext.prefix, MAX_COMPLETION_TABLES, completionContext.qualifier, props.catalog);
let schemaTables = connectionStore.lookupLocalCompletionTables(props.connectionId!, scope.database, completionContext.prefix, MAX_COMPLETION_TABLES, completionContext.qualifier, props.catalog);
if (!localOnlyMetadata) {
schemaTables = await listCompletionTablesWithLatencyBudget(props.connectionId!, props.database!, completionContext.prefix, MAX_COMPLETION_TABLES, completionContext.qualifier);
schemaTables = await listCompletionTablesWithLatencyBudget(props.connectionId!, scope.database, completionContext.prefix, MAX_COMPLETION_TABLES, completionContext.qualifier, false, props.catalog, scope.schema);
} else if (schemaTables.length === 0 && allowsOnDemandQualifiedTableCompletion(completionContext.prefix)) {
schemaTables = await listCompletionTablesWithLatencyBudget(props.connectionId!, props.database!, completionContext.prefix, PRESTO_ON_DEMAND_TABLE_COMPLETION_LIMIT, completionContext.qualifier);
schemaTables = await listCompletionTablesWithLatencyBudget(props.connectionId!, scope.database, completionContext.prefix, PRESTO_ON_DEMAND_TABLE_COMPLETION_LIMIT, completionContext.qualifier, false, props.catalog, scope.schema);
}
if (schemaTables.length > 0) {
tables = schemaTables;
@ -3228,7 +3343,12 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
});
const unresolvedRefs = refs.filter((rt) => !usesOracleSessionCompletionColumns(rt.schema) && !rt.schema && !rt.columns && !isVirtualCompletionTableReference(rt));
if (!localOnlyMetadata && unresolvedRefs.length > 0) {
const lookupGroups = await Promise.all(unresolvedRefs.map((rt) => connectionStore.listCompletionTables(props.connectionId!, props.database!, rt.name, 20, props.schema, false, props.schema, props.catalog)));
const lookupGroups = await Promise.all(
unresolvedRefs.map((rt) => {
const target = completionMetadataTarget(rt, scope);
return connectionStore.listCompletionTables(props.connectionId!, target?.database ?? scope.database, rt.name, 20, target?.schema ?? scope.schema, false, scope.schema, target?.catalog ?? props.catalog);
}),
);
if (epoch !== completionEpoch) return null;
const lookupTables = lookupGroups.flat();
refs = refs.map((rt) => {
@ -3268,10 +3388,10 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
refs.map(async (refTable) => {
if (isVirtualCompletionTableReference(refTable)) return;
if (refTable.columns && refTable.columns.length > 0) return;
const cacheKey = completionCacheKey(refTable);
const cacheKey = completionCacheKey(refTable, scope);
if (cachedColumnsByTable.has(cacheKey)) return;
try {
const target = completionMetadataTarget(refTable);
const target = completionMetadataTarget(refTable, scope);
if (!target) return;
const columns = await listCompletionColumnsForEditor(props.connectionId!, target.database, refTable.name, target.schema, target.catalog, refTable);
if (epoch !== completionEpoch) return;
@ -3312,14 +3432,14 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
);
continue;
}
const cacheKey = completionCacheKey(refTable);
const cacheKey = completionCacheKey(refTable, scope);
const cached = cachedColumnsByTable.get(cacheKey);
if (cached) {
columnsByTable.set(cacheKey, cached);
}
let cachedForeignKeys = cachedForeignKeysByTable.get(cacheKey);
if (!cachedForeignKeys) {
const target = completionMetadataTarget(refTable);
const target = completionMetadataTarget(refTable, scope);
cachedForeignKeys = target ? connectionStore.lookupLocalCompletionForeignKeys(props.connectionId!, target.database, refTable.name, target.schema) : [];
if (cachedForeignKeys.length > 0) cachedForeignKeysByTable.set(cacheKey, cachedForeignKeys);
}
@ -3349,7 +3469,7 @@ async function performAsyncCompletionWithResult(epoch: number, completionContext
snippets: settingsStore.editorSettings.snippets,
dialect: props.dialect,
databaseType: snippetDatabaseType.value,
currentSchema: props.schema,
currentSchema: scope.schema,
keywordCase: settingsStore.editorSettings.sqlFormatter.keywordCase,
autoAliasTables: settingsStore.editorSettings.autoAliasTables,
});
@ -3384,6 +3504,14 @@ function mergeCompletionObjects(existing: SqlCompletionObject[], incoming: SqlCo
return merged;
}
function completionObjectScopeKey(scope: CompletionMetadataScope): string {
return `${scope.database}:${scope.schema ?? ""}`.toLowerCase();
}
function completionObjectsForScope(scope: CompletionMetadataScope): SqlCompletionObject[] {
return cachedCompletionObjectsByScope.get(completionObjectScopeKey(scope)) ?? [];
}
function completionObjectIdentityKey(object: SqlCompletionObject): string {
return `${object.type}:${object.schema ?? ""}:${object.name}:${object.parentName ?? ""}:${object.signature?.trim() ?? ""}`.toLowerCase();
}
@ -3419,7 +3547,7 @@ function refreshActiveSqlCompletion(fullDoc: string, position: number, completio
function refreshCompletionCache() {
cachedTables = [];
cachedCompletionObjects = [];
cachedCompletionObjectsByScope.clear();
cachedColumnsByTable.clear();
cachedInsertValueHintColumnsByTable.clear();
loadedColumnsByTable.clear();
@ -4217,7 +4345,7 @@ onMounted(async () => {
registerTableReferenceDropListener();
cachedTables = [];
cachedCompletionObjects = [];
cachedCompletionObjectsByScope.clear();
scheduleSemanticDiagnostics();
if (props.autoFocus) {

View File

@ -34,4 +34,23 @@ describe("QueryEditor database name completion wiring", () => {
expect(guard).toContain("supportsDatabaseNameCompletion(props.databaseType)");
expect(guard).toContain("supportsDatabaseSchemaQualifierCompletion()");
});
it("uses the resolved database scope for routine completion and isolates the editor cache", () => {
expect(extractFunction("routineCompletionTargetForContext")).toContain("currentDatabase: scope.database");
expect(extractFunction("routineCompletionTargetForContext")).toContain("supportsDatabaseSchemaQualifier: supportsDatabaseSchemaQualifierCompletion()");
expect(extractFunction("lookupLocalCompletionObjectsForContext")).toContain("target.database");
expect(extractFunction("listCompletionObjectsForContext")).toContain("target.database");
expect(extractFunction("routineCompletionScopeForContext")).toContain("database: target.database");
expect(extractFunction("completionObjectScopeKey")).toContain("scope.database");
expect(extractFunction("completionObjectScopeKey")).toContain("scope.schema");
expect(queryEditorSource).toContain("cachedCompletionObjectsByScope");
});
it("loads SQL Server capability metadata before offering or applying USE completion", () => {
const provider = extractFunction("provideSqlCompletions");
expect(provider).toContain("getSqlServerCompletionContext");
expect(provider).toContain("supports_session_database_switch");
expect(provider).toContain("useDatabaseDefaultSchema");
});
});

View File

@ -1,6 +1,15 @@
import { describe, expect, it } from "vitest";
import { getSqlCompletionContext } from "@/lib/sql/sqlCompletion";
import { mergeSqlCompletionQualifierNames, resolveSqlCompletionRoutineLookupTarget, resolveSqlCompletionSchemaLookupDatabase, resolveSqlCompletionTableLookupTarget } from "@/lib/sql/sqlCompletionLookupTarget";
import {
buildSqlServerUseDatabaseCompletionItems,
mergeSqlCompletionQualifierNames,
resolveSqlCompletionRoutineLookupTarget,
resolveSqlCompletionSchemaLookupDatabase,
resolveSqlCompletionScope,
resolveSqlCompletionTableLookupTarget,
resolveSqlServerUseDatabaseCompletion,
sqlServerUseCompletionDatabaseNames,
} from "@/lib/sql/sqlCompletionLookupTarget";
describe("sqlCompletionLookupTarget", () => {
it("treats qualified table completion as a database lookup for MySQL-compatible engines", () => {
@ -215,22 +224,252 @@ describe("sqlCompletionLookupTarget", () => {
});
});
it("uses the preceding SQL Server USE database and its server-reported default schema", () => {
const sql = "USE [bardb]\n\nSELECT * FROM T";
const completionContext = getSqlCompletionContext(sql, sql.length, { databaseType: "sqlserver", dialect: "sqlserver" });
const scope = resolveSqlCompletionScope({
sql,
cursor: sql.length,
databaseType: "sqlserver",
currentDatabase: "FooDB",
currentSchema: "sales",
knownDatabases: ["FooDB", "BarDB"],
supportsSessionDatabaseSwitch: true,
useDatabaseDefaultSchema: "app_user",
completionContext,
});
expect(scope.database).toBe("BarDB");
expect(scope.schema).toBe("app_user");
expect(
resolveSqlCompletionTableLookupTarget({
currentDatabase: scope.database,
currentSchema: scope.schema,
supportsDatabaseQualifier: false,
supportsDatabaseSchemaQualifier: true,
completionContext: scope.completionContext,
}),
).toEqual({
database: "BarDB",
schema: "app_user",
filter: "T",
});
});
it("uses the last preceding SQL Server USE and unescapes its identifier", () => {
const sql = "USE FooDB;\nUSE [Bar]]DB];\nSELECT * FROM T";
const completionContext = getSqlCompletionContext(sql, sql.length, { databaseType: "sqlserver", dialect: "sqlserver" });
const scope = resolveSqlCompletionScope({
sql,
cursor: sql.length,
databaseType: "sqlserver",
currentDatabase: "SelectedDB",
knownDatabases: ["SelectedDB", "FooDB", "Bar]DB"],
supportsSessionDatabaseSwitch: true,
useDatabaseDefaultSchema: "reporting_user",
completionContext,
});
expect(scope.database).toBe("Bar]DB");
});
it("ignores commented, quoted, current, later, and non-SQL Server USE text", () => {
const sql = "-- USE [CommentDB]\nSELECT 'USE [StringDB]';\nSELECT * FROM T;\nUSE [LaterDB];";
const cursor = sql.indexOf("T;") + 1;
const completionContext = getSqlCompletionContext(sql, cursor, { databaseType: "sqlserver", dialect: "sqlserver" });
expect(
resolveSqlCompletionScope({
sql,
cursor,
databaseType: "sqlserver",
currentDatabase: "FooDB",
currentSchema: "dbo",
completionContext,
}).database,
).toBe("FooDB");
const currentUseSql = "USE [BarDB]";
expect(
resolveSqlCompletionScope({
sql: currentUseSql,
cursor: currentUseSql.length,
databaseType: "sqlserver",
currentDatabase: "FooDB",
currentSchema: "dbo",
completionContext,
}).database,
).toBe("FooDB");
expect(
resolveSqlCompletionScope({
sql: "USE [BarDB];\nSELECT * FROM T",
cursor: "USE [BarDB];\nSELECT * FROM T".length,
databaseType: "postgres",
currentDatabase: "FooDB",
currentSchema: "public",
completionContext,
}),
).toMatchObject({ database: "FooDB", schema: "public", completionContext });
});
it("scopes unqualified SQL Server references without overriding explicit databases or schemas", () => {
const sql = "USE [BarDB];\nSELECT * FROM T";
const completionContext = {
...getSqlCompletionContext(sql, sql.length, { databaseType: "sqlserver", dialect: "sqlserver" }),
insertTable: "NewRow",
referencedTables: [{ name: "TUser" }, { name: "TOrder", schema: "sales" }, { name: "TArchive", database: "ArchiveDB", schema: "history" }],
};
const scope = resolveSqlCompletionScope({
sql,
cursor: sql.length,
databaseType: "sqlserver",
currentDatabase: "FooDB",
knownDatabases: ["FooDB", "BarDB", "ArchiveDB"],
supportsSessionDatabaseSwitch: true,
useDatabaseDefaultSchema: "app_user",
completionContext,
});
expect(scope.completionContext).toMatchObject({
insertDatabase: "BarDB",
insertSchema: "app_user",
referencedTables: [
{ name: "TUser", database: "BarDB", schema: "app_user" },
{ name: "TOrder", database: "BarDB", schema: "sales" },
{ name: "TArchive", database: "ArchiveDB", schema: "history" },
],
});
});
it("falls back to the selected SQL Server database when USE names an unknown database", () => {
const sql = "USE [MissingDB];\nSELECT * FROM T";
const completionContext = getSqlCompletionContext(sql, sql.length, { databaseType: "sqlserver", dialect: "sqlserver" });
const scope = resolveSqlCompletionScope({
sql,
cursor: sql.length,
databaseType: "sqlserver",
currentDatabase: "FooDB",
currentSchema: "sales",
knownDatabases: ["FooDB", "BarDB"],
supportsSessionDatabaseSwitch: true,
useDatabaseDefaultSchema: "dbo",
completionContext,
});
expect(scope).toEqual({
database: "FooDB",
schema: "sales",
completionContext,
});
});
it("falls back to the selected database when the endpoint cannot switch sessions", () => {
const sql = "USE [BarDB];\nSELECT * FROM T";
const completionContext = getSqlCompletionContext(sql, sql.length, { databaseType: "sqlserver", dialect: "sqlserver" });
expect(
resolveSqlCompletionScope({
sql,
cursor: sql.length,
databaseType: "sqlserver",
currentDatabase: "AzureDB",
currentSchema: "sales",
knownDatabases: ["AzureDB", "BarDB"],
supportsSessionDatabaseSwitch: false,
useDatabaseDefaultSchema: "dbo",
completionContext,
}),
).toEqual({
database: "AzureDB",
schema: "sales",
completionContext,
});
});
it.each([
["USE ", { from: 4, prefix: "", quoteStyle: "none" }],
["USE Bar", { from: 4, prefix: "Bar", quoteStyle: "none" }],
["USE [Bar", { from: 5, prefix: "Bar", quoteStyle: "bracket" }],
['USE "Bar', { from: 5, prefix: "Bar", quoteStyle: "double" }],
["SELECT 1;\nUSE [Bar", { from: 15, prefix: "Bar", quoteStyle: "bracket" }],
])("resolves SQL Server database completion for %s", (sql, expected) => {
expect(resolveSqlServerUseDatabaseCompletion({ sql, cursor: sql.length, databaseType: "sqlserver" })).toEqual(expected);
});
it("does not offer USE database completion outside an incomplete SQL Server USE target", () => {
expect(resolveSqlServerUseDatabaseCompletion({ sql: "USE", cursor: 3, databaseType: "sqlserver" })).toBeUndefined();
expect(resolveSqlServerUseDatabaseCompletion({ sql: "USE [BarDB]", cursor: 11, databaseType: "sqlserver" })).toBeUndefined();
expect(resolveSqlServerUseDatabaseCompletion({ sql: "SELECT 'USE Bar'", cursor: 15, databaseType: "sqlserver" })).toBeUndefined();
expect(resolveSqlServerUseDatabaseCompletion({ sql: "USE Bar", cursor: 7, databaseType: "postgres" })).toBeUndefined();
});
it("builds quoted SQL Server USE database completion insertions", () => {
const unquoted = resolveSqlServerUseDatabaseCompletion({ sql: "USE Odd", cursor: 7, databaseType: "sqlserver" })!;
const bracketed = resolveSqlServerUseDatabaseCompletion({ sql: "USE [Odd", cursor: 8, databaseType: "sqlserver" })!;
const doubleQuoted = resolveSqlServerUseDatabaseCompletion({ sql: 'USE "Odd', cursor: 8, databaseType: "sqlserver" })!;
expect(buildSqlServerUseDatabaseCompletionItems(["Odd]DB"], unquoted)[0]).toMatchObject({ label: "Odd]DB", detail: "database", apply: "[Odd]]DB]" });
expect(buildSqlServerUseDatabaseCompletionItems(["Odd]DB"], bracketed)[0]).toMatchObject({ label: "Odd]DB", filterText: "Odd]]DB", apply: "Odd]]DB]" });
expect(buildSqlServerUseDatabaseCompletionItems(['Odd"DB'], doubleQuoted)[0]).toMatchObject({ label: 'Odd"DB', filterText: 'Odd""DB', apply: 'Odd""DB"' });
});
it("limits USE database candidates to the current database when session switching is unsupported", () => {
expect(
sqlServerUseCompletionDatabaseNames({
databaseNames: ["master", "AzureDB", "OtherDB"],
currentDatabase: "azuredb",
supportsSessionDatabaseSwitch: false,
}),
).toEqual(["AzureDB"]);
expect(
sqlServerUseCompletionDatabaseNames({
databaseNames: ["FooDB", "BarDB"],
currentDatabase: "FooDB",
supportsSessionDatabaseSwitch: true,
}),
).toEqual(["FooDB", "BarDB"]);
expect(
sqlServerUseCompletionDatabaseNames({
databaseNames: [],
currentDatabase: "",
supportsSessionDatabaseSwitch: false,
}),
).toEqual([]);
});
it.each([
["SELECT dbo.fn_", "dbo", "fn_"],
["SELECT public.st_", "public", "st_"],
])("separates a routine schema from its name mask for %s", (sql, schema, mask) => {
const completionContext = getSqlCompletionContext(sql, sql.length);
expect(resolveSqlCompletionRoutineLookupTarget({ currentSchema: "fallback", completionContext })).toEqual({ schema, mask });
expect(resolveSqlCompletionRoutineLookupTarget({ currentDatabase: "app", currentSchema: "fallback", completionContext })).toEqual({ database: "app", schema, mask });
});
it("uses the current schema for an unqualified routine mask", () => {
const sql = "SELECT st_";
const completionContext = getSqlCompletionContext(sql, sql.length);
expect(resolveSqlCompletionRoutineLookupTarget({ currentSchema: "public", completionContext })).toEqual({
expect(resolveSqlCompletionRoutineLookupTarget({ currentDatabase: "app", currentSchema: "public", completionContext })).toEqual({
database: "app",
schema: "public",
mask: "st_",
});
});
it.each(["SELECT BarDB.sales.fn_", "EXEC BarDB.sales.proc_"])("uses an explicit SQL Server database and schema for %s", (sql) => {
const completionContext = getSqlCompletionContext(sql, sql.length);
expect(
resolveSqlCompletionRoutineLookupTarget({
currentDatabase: "FooDB",
currentSchema: "app_user",
supportsDatabaseSchemaQualifier: true,
completionContext,
}),
).toEqual({
database: "BarDB",
schema: "sales",
mask: sql.endsWith("fn_") ? "fn_" : "proc_",
});
});
});

View File

@ -138,6 +138,7 @@ export const syncSavedSqlDirectory = forward("syncSavedSqlDirectory");
// Schema
export const listDatabases = forward("listDatabases");
export const listDatabaseStorage = forward("listDatabaseStorage");
export const getSqlServerCompletionContext = forward("getSqlServerCompletionContext");
export const listDorisCatalogs = forward("listDorisCatalogs");
export const listDorisCatalogDatabases = forward("listDorisCatalogDatabases");
export const listSqlServerLinkedServers = forward("listSqlServerLinkedServers");

View File

@ -4,6 +4,7 @@ import type {
DatabaseConnectionInfo,
DatabaseInfo,
DatabaseStorageInfo,
SqlServerCompletionContext,
SchemaInfo,
LinkedServerInfo,
CatalogInfo,
@ -634,6 +635,10 @@ export async function listDatabaseStorage(connectionId: string, databases: strin
});
}
export async function getSqlServerCompletionContext(connectionId: string, database: string): Promise<SqlServerCompletionContext> {
return get(`/api/schema/sqlserver/completion-context?${qs({ connection_id: connectionId, database })}`);
}
export async function listDorisCatalogs(connectionId: string): Promise<CatalogInfo[]> {
return get(`/api/schema/doris/catalogs?${qs({ connection_id: connectionId })}`);
}

View File

@ -8,6 +8,7 @@ import type {
DatabaseConnectionInfo,
DatabaseInfo,
DatabaseStorageInfo,
SqlServerCompletionContext,
SchemaInfo,
LinkedServerInfo,
CatalogInfo,
@ -869,6 +870,10 @@ export async function listDatabaseStorage(connectionId: string, databases: strin
return invoke("list_database_storage", { connectionId, databases });
}
export async function getSqlServerCompletionContext(connectionId: string, database: string): Promise<SqlServerCompletionContext> {
return invoke("get_sqlserver_completion_context", { connectionId, database });
}
export async function listDorisCatalogs(connectionId: string): Promise<CatalogInfo[]> {
return invoke("list_doris_catalogs", { connectionId });
}

View File

@ -1,4 +1,6 @@
import type { SqlCompletionContext } from "@/lib/sql/sqlCompletion";
import type { SqlCompletionContext, SqlCompletionItem } from "@/lib/sql/sqlCompletion";
import { currentExecutableStatementRange, executableStatementRanges } from "@/lib/sql/sqlStatementRanges";
import type { DatabaseType } from "@/types/database";
export interface SqlCompletionTableLookupTarget {
database: string;
@ -8,10 +10,186 @@ export interface SqlCompletionTableLookupTarget {
}
export interface SqlCompletionRoutineLookupTarget {
database: string;
schema?: string;
mask: string;
}
export interface SqlCompletionScope {
database: string;
schema?: string;
completionContext: SqlCompletionContext;
}
export interface SqlServerUseDatabaseCompletion {
from: number;
prefix: string;
quoteStyle: "none" | "bracket" | "double";
}
function sqlStatementWithoutLeadingComments(statement: string): string {
let remaining = statement.trimStart();
while (remaining) {
if (remaining.startsWith("--")) {
const newline = remaining.indexOf("\n");
remaining = newline < 0 ? "" : remaining.slice(newline + 1).trimStart();
continue;
}
if (remaining.startsWith("/*")) {
const end = remaining.indexOf("*/", 2);
if (end < 0) return "";
remaining = remaining.slice(end + 2).trimStart();
continue;
}
break;
}
return remaining;
}
function sqlServerUseDatabase(statement: string): string | undefined {
const match = /^USE\s+(?:\[((?:[^\]]|\]\])*)\]|"((?:[^"]|"")*)"|([\p{L}_@#][\p{L}\p{N}_@$#]*))\s*;?\s*$/iu.exec(sqlStatementWithoutLeadingComments(statement));
if (!match) return undefined;
if (match[1] !== undefined) return match[1].replaceAll("]]", "]");
if (match[2] !== undefined) return match[2].replaceAll('""', '"');
return match[3];
}
export function sqlServerUseDatabaseBeforeCursor(sql: string, cursor: number): string | undefined {
const position = Math.max(0, Math.min(cursor, sql.length));
let database: string | undefined;
for (const statement of executableStatementRanges(sql, "sqlserver")) {
if (statement.from >= position || statement.to >= position) break;
database = sqlServerUseDatabase(statement.sql) ?? database;
}
return database;
}
function unclosedQuotedIdentifierPrefix(value: string, quoteStyle: "bracket" | "double"): string | undefined {
const closingQuote = quoteStyle === "bracket" ? "]" : '"';
let prefix = "";
for (let index = 1; index < value.length; index += 1) {
const character = value[index]!;
if (character !== closingQuote) {
prefix += character;
continue;
}
if (value[index + 1] !== closingQuote) return undefined;
prefix += closingQuote;
index += 1;
}
return prefix;
}
export function resolveSqlServerUseDatabaseCompletion(options: { sql: string; cursor: number; databaseType?: DatabaseType }): SqlServerUseDatabaseCompletion | undefined {
if (options.databaseType !== "sqlserver") return undefined;
const position = Math.max(0, Math.min(options.cursor, options.sql.length));
const statement = currentExecutableStatementRange(options.sql, position, "sqlserver");
if (!statement || (statement.to > position && options.sql.slice(position, statement.to).trim())) return undefined;
const beforeCursor = options.sql.slice(statement.from, position);
const useMatch = /^USE(?=\s)/iu.exec(beforeCursor);
if (!useMatch) return undefined;
let targetOffset = useMatch[0].length;
while (targetOffset < beforeCursor.length && /\s/u.test(beforeCursor[targetOffset]!)) targetOffset += 1;
const target = beforeCursor.slice(targetOffset);
if (!target) {
return {
from: statement.from + targetOffset,
prefix: "",
quoteStyle: "none",
};
}
if (/^[\p{L}_@#][\p{L}\p{N}_@$#]*$/u.test(target)) {
return {
from: statement.from + targetOffset,
prefix: target,
quoteStyle: "none",
};
}
const quoteStyle = target[0] === "[" ? "bracket" : target[0] === '"' ? "double" : undefined;
if (!quoteStyle) return undefined;
const prefix = unclosedQuotedIdentifierPrefix(target, quoteStyle);
if (prefix === undefined) return undefined;
return {
from: statement.from + targetOffset + 1,
prefix,
quoteStyle,
};
}
export function buildSqlServerUseDatabaseCompletionItems(databaseNames: readonly string[], completion: SqlServerUseDatabaseCompletion): SqlCompletionItem[] {
return databaseNames.map((database) => {
const escapedDatabase = completion.quoteStyle === "double" ? database.replaceAll('"', '""') : database.replaceAll("]", "]]");
const apply = completion.quoteStyle === "bracket" ? `${escapedDatabase}]` : completion.quoteStyle === "double" ? `${escapedDatabase}"` : `[${escapedDatabase}]`;
return {
label: database,
filterText: completion.quoteStyle === "none" ? database : escapedDatabase,
type: "schema",
detail: "database",
apply,
boost: 1_500,
};
});
}
export function sqlServerUseCompletionDatabaseNames(options: { databaseNames: readonly string[]; currentDatabase: string; supportsSessionDatabaseSwitch: boolean }): string[] {
if (options.supportsSessionDatabaseSwitch) return [...options.databaseNames];
const currentDatabase = options.currentDatabase.trim();
return currentDatabase ? [findExactName(options.databaseNames, currentDatabase) ?? currentDatabase] : [];
}
export function resolveSqlCompletionScope(options: {
sql: string;
cursor: number;
databaseType?: DatabaseType;
currentDatabase: string;
currentSchema?: string;
knownDatabases?: readonly string[];
supportsSessionDatabaseSwitch?: boolean;
useDatabaseDefaultSchema?: string;
completionContext: SqlCompletionContext;
}): SqlCompletionScope {
if (options.databaseType !== "sqlserver") {
return {
database: options.currentDatabase,
schema: options.currentSchema,
completionContext: options.completionContext,
};
}
const parsedDatabase = sqlServerUseDatabaseBeforeCursor(options.sql, options.cursor);
const database = parsedDatabase ? findExactName(options.knownDatabases, parsedDatabase) : undefined;
const targetsCurrentDatabase = database?.toLowerCase() === options.currentDatabase.toLowerCase();
const schema = options.useDatabaseDefaultSchema?.trim();
if (!database || !schema || (!targetsCurrentDatabase && options.supportsSessionDatabaseSwitch !== true)) {
return {
database: options.currentDatabase,
schema: options.currentSchema,
completionContext: options.completionContext,
};
}
return {
database,
schema,
completionContext: {
...options.completionContext,
insertDatabase: options.completionContext.insertTable && !options.completionContext.insertDatabase ? database : options.completionContext.insertDatabase,
insertSchema: options.completionContext.insertTable && !options.completionContext.insertSchema ? schema : options.completionContext.insertSchema,
referencedTables: options.completionContext.referencedTables.map((table) =>
table.database
? table
: {
...table,
database,
schema: table.schema ?? schema,
},
),
},
};
}
function findExactName(names: readonly string[] | undefined, value: string): string | undefined {
return names?.find((name) => name.toLowerCase() === value.toLowerCase());
}
@ -82,13 +260,17 @@ export function resolveSqlCompletionTableLookupTarget(options: {
};
}
export function resolveSqlCompletionRoutineLookupTarget(options: { currentSchema?: string; completionContext: Pick<SqlCompletionContext, "qualifier" | "qualifierParts" | "prefix"> }): SqlCompletionRoutineLookupTarget {
const qualifierParts = options.completionContext.qualifierParts?.filter(Boolean);
const schema = qualifierParts?.[qualifierParts.length - 1] ?? options.completionContext.qualifier?.trim() ?? options.currentSchema;
export function resolveSqlCompletionRoutineLookupTarget(options: { currentDatabase: string; currentSchema?: string; supportsDatabaseSchemaQualifier?: boolean; completionContext: Pick<SqlCompletionContext, "qualifier" | "qualifierParts" | "prefix"> }): SqlCompletionRoutineLookupTarget {
const qualifier = options.completionContext.qualifier?.trim();
const qualifierParts = options.completionContext.qualifierParts?.filter(Boolean) ?? qualifier?.split(".").filter(Boolean) ?? [];
const hasDatabaseQualifier = options.supportsDatabaseSchemaQualifier && qualifierParts.length >= 2;
const database = hasDatabaseQualifier ? qualifierParts[qualifierParts.length - 2]! : options.currentDatabase;
const schema = qualifierParts[qualifierParts.length - 1] ?? qualifier ?? options.currentSchema;
// A qualified routine uses the qualifier as metadata scope; only the final
// identifier fragment is the function/procedure name mask.
return {
database,
schema: schema || undefined,
mask: options.completionContext.prefix,
};

View File

@ -557,9 +557,9 @@ describe("connectionStore completion assistant", () => {
]);
});
it("searches default SQL Server schemas without treating the username as a schema", async () => {
it("searches the server-reported SQL Server default schema without treating the username as a schema", async () => {
const completionAssistantSearch = vi.fn().mockResolvedValue({
candidates: [{ name: "st_area", kind: "function", schema: "dbo", data_type: "float" }],
candidates: [{ name: "st_area", kind: "function", schema: "app_user", data_type: "float" }],
incomplete: false,
fallback_used: false,
});
@ -576,11 +576,11 @@ describe("connectionStore completion assistant", () => {
store.connections = [sqlServerConnection()];
store.connectedIds.add("sqlserver-1");
const objects = await store.listCompletionObjects("sqlserver-1", "app", "st_", 20);
const objects = await store.listCompletionObjects("sqlserver-1", "app", "st_", 20, undefined, undefined, false, "app_user");
expect(completionAssistantSearch).toHaveBeenCalledWith(
expect.objectContaining({
schema: null,
schema: "app_user",
parent_schema: null,
mask: "st_",
}),
@ -588,15 +588,100 @@ describe("connectionStore completion assistant", () => {
expect(objects).toEqual([
expect.objectContaining({
name: "st_area",
schema: "dbo",
schema: "app_user",
type: "function",
dataType: "float",
applyName: "dbo.st_area",
applyName: "app_user.st_area",
boost: 1000,
}),
]);
});
it("prefers an explicit SQL Server routine schema over the current default schema", async () => {
const completionAssistantSearch = vi.fn().mockResolvedValue({
candidates: [{ name: "calculate_tax", kind: "function", schema: "sales", data_type: "decimal" }],
incomplete: false,
fallback_used: false,
});
vi.doMock("@/lib/backend/tauriRuntime", () => ({ isTauriRuntime: () => false }));
vi.doMock("@/lib/backend/api", () => ({
checkConnectionHealth: vi.fn().mockResolvedValue(undefined),
completionAssistantSearch,
listCompletionObjects: vi.fn().mockResolvedValue([]),
}));
const { useConnectionStore } = await import("@/stores/connectionStore");
const store = useConnectionStore();
store.connections = [sqlServerConnection()];
store.connectedIds.add("sqlserver-1");
await store.listCompletionObjects("sqlserver-1", "BarDB", "calculate_", 20, "sales", undefined, false, "app_user");
expect(completionAssistantSearch).toHaveBeenCalledWith(
expect.objectContaining({
database: "BarDB",
schema: "sales",
parent_schema: "sales",
mask: "calculate_",
}),
);
});
it("caches SQL Server completion context independently for each database", async () => {
const getSqlServerCompletionContext = vi.fn(async (_connectionId: string, database: string) => ({
default_schema: database === "BarDB" ? "bar_user" : "foo_user",
supports_session_database_switch: true,
}));
vi.doMock("@/lib/backend/tauriRuntime", () => ({ isTauriRuntime: () => false }));
vi.doMock("@/lib/backend/api", () => ({
checkConnectionHealth: vi.fn().mockResolvedValue(undefined),
getSqlServerCompletionContext,
}));
const { useConnectionStore } = await import("@/stores/connectionStore");
const store = useConnectionStore();
store.connections = [sqlServerConnection()];
store.connectedIds.add("sqlserver-1");
const foo = await store.getSqlServerCompletionContext("sqlserver-1", "FooDB");
const fooCached = await store.getSqlServerCompletionContext("sqlserver-1", "FooDB");
const bar = await store.getSqlServerCompletionContext("sqlserver-1", "BarDB");
expect(foo).toMatchObject({ default_schema: "foo_user" });
expect(fooCached).toEqual(foo);
expect(bar).toMatchObject({ default_schema: "bar_user" });
expect(getSqlServerCompletionContext).toHaveBeenCalledTimes(2);
});
it("keeps SQL Server routine results isolated across databases", async () => {
const completionAssistantSearch = vi.fn(async (request: { database: string }) => ({
candidates: [{ name: request.database === "BarDB" ? "bar_proc" : "foo_proc", kind: "procedure", schema: "app_user" }],
incomplete: false,
fallback_used: false,
}));
vi.doMock("@/lib/backend/tauriRuntime", () => ({ isTauriRuntime: () => false }));
vi.doMock("@/lib/backend/api", () => ({
checkConnectionHealth: vi.fn().mockResolvedValue(undefined),
completionAssistantSearch,
listCompletionObjects: vi.fn().mockResolvedValue([]),
}));
const { useConnectionStore } = await import("@/stores/connectionStore");
const store = useConnectionStore();
store.connections = [sqlServerConnection()];
store.connectedIds.add("sqlserver-1");
const foo = await store.listCompletionObjects("sqlserver-1", "FooDB", "", 20, undefined, undefined, false, "app_user");
const bar = await store.listCompletionObjects("sqlserver-1", "BarDB", "", 20, undefined, undefined, false, "app_user");
expect(foo.map((object) => object.name)).toEqual(["foo_proc"]);
expect(bar.map((object) => object.name)).toEqual(["bar_proc"]);
expect(completionAssistantSearch).toHaveBeenCalledTimes(2);
});
it("limits concurrent completion column metadata requests per connection database", async () => {
const gates = [deferred<any[]>(), deferred<any[]>(), deferred<any[]>(), deferred<any[]>()];
let activeColumns = 0;

View File

@ -10,6 +10,7 @@ import type {
DatabaseType,
DatabaseConnectionInfo,
DatabaseStorageInfo,
SqlServerCompletionContext,
CatalogInfo,
ForeignKeyInfo,
ObjectInfo,
@ -340,6 +341,7 @@ export const useConnectionStore = defineStore("connection", () => {
const completionColumnsCache = ref<Record<string, ColumnInfo[]>>({});
const completionForeignKeysCache = ref<Record<string, ForeignKeyInfo[]>>({});
const completionDatabasesCache = ref<Record<string, string[]>>({});
const sqlServerCompletionContextCache = ref<Record<string, SqlServerCompletionContext>>({});
const elasticsearchCompletionIndicesCache = ref<Record<string, string[]>>({});
const redisCompletionKeysCache = ref<Record<string, string[]>>({});
const mongoCompletionCollectionsCache = ref<Record<string, string[]>>({});
@ -2310,6 +2312,9 @@ export const useConnectionStore = defineStore("connection", () => {
for (const key of Object.keys(schemaListCache.value)) {
if (key === exactCacheKey || key.startsWith(cachePrefix)) delete schemaListCache.value[key];
}
for (const key of Object.keys(sqlServerCompletionContextCache.value)) {
if (key === exactCacheKey || key.startsWith(cachePrefix)) delete sqlServerCompletionContextCache.value[key];
}
for (const key of Object.keys(elasticsearchCompletionIndicesCache.value)) {
if (key === exactCacheKey || key.startsWith(cachePrefix)) delete elasticsearchCompletionIndicesCache.value[key];
}
@ -5417,8 +5422,8 @@ export const useConnectionStore = defineStore("connection", () => {
): Promise<SqlCompletionObject[]> {
const databaseType = getConfig(connectionId)?.db_type;
const oracleAssistant = databaseType === "oracle";
const requestedSchema = currentSchema?.trim() || schema?.trim() || undefined;
const preferredSchema = oracleAssistant ? completionPreferredSchema(connectionId, currentSchema) : requestedSchema || (databaseType === "sqlserver" ? "dbo" : databaseType === "postgres" ? "public" : databaseType === "mysql" ? database : undefined);
const requestedSchema = schema?.trim() || currentSchema?.trim() || undefined;
const preferredSchema = oracleAssistant ? completionPreferredSchema(connectionId, currentSchema) : requestedSchema || (databaseType === "postgres" ? "public" : databaseType === "mysql" ? database : undefined);
const response = await completionAssistantSearch({
connection_id: connectionId,
database,
@ -5666,6 +5671,20 @@ export const useConnectionStore = defineStore("connection", () => {
);
}
async function getSqlServerCompletionContext(connectionId: string, database: string): Promise<SqlServerCompletionContext> {
const cacheKey = `${connectionId}:${database}`;
if (sqlServerCompletionContextCache.value[cacheKey]) {
return sqlServerCompletionContextCache.value[cacheKey];
}
return withCompletionInFlight(`${cacheKey}:sqlserver-completion-context`, async () => {
await ensureConnected(connectionId);
const context = await api.getSqlServerCompletionContext(connectionId, database);
sqlServerCompletionContextCache.value[cacheKey] = context;
evictOldestCacheEntries(sqlServerCompletionContextCache.value, COMPLETION_CACHE_MAX);
return context;
});
}
async function listCompletionSchemas(connectionId: string, database: string): Promise<string[]> {
const cacheKey = `${connectionId}:${database}`;
if (schemaListCache.value[cacheKey]) {
@ -6776,6 +6795,7 @@ export const useConnectionStore = defineStore("connection", () => {
listCompletionForeignKeys,
listCompletionSchemas,
listCompletionDatabases,
getSqlServerCompletionContext,
lookupLocalCompletionTables,
lookupLocalCompletionObjects,
lookupLocalCompletionColumns,

View File

@ -370,6 +370,11 @@ export interface DatabaseStorageInfo {
size_bytes: number | null;
}
export interface SqlServerCompletionContext {
default_schema: string;
supports_session_database_switch: boolean;
}
export interface SchemaInfo {
name: string;
comment?: string | null;

View File

@ -92,6 +92,56 @@ pub struct SqlServerColumnMetadata {
pub generated_always_type: i32,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
pub struct SqlServerCompletionContext {
pub default_schema: String,
pub supports_session_database_switch: bool,
}
const SQLSERVER_COMPLETION_CONTEXT_SQL: &str = "\
SELECT COALESCE(\
(SELECT default_schema.name \
FROM sys.schemas default_schema \
WHERE default_schema.name = SCHEMA_NAME()), \
N'dbo'\
) AS default_schema, \
CONVERT(int, SERVERPROPERTY(N'EngineEdition')) AS engine_edition";
fn sqlserver_supports_session_database_switch(engine_edition: i32) -> bool {
// Only known boxed SQL Server, Managed Instance, and SQL Edge editions are
// allowed to switch databases. Cloud single-database endpoints and future
// editions default to opening a connection directly to the target database.
matches!(engine_edition, 1 | 2 | 3 | 4 | 8 | 9)
}
fn sqlserver_completion_context(
default_schema: Option<&str>,
engine_edition: Option<i32>,
) -> Result<SqlServerCompletionContext, String> {
let default_schema = default_schema.map(str::trim).filter(|schema| !schema.is_empty()).unwrap_or("dbo");
let engine_edition = engine_edition.ok_or_else(|| "SQL Server EngineEdition is unavailable".to_string())?;
Ok(SqlServerCompletionContext {
default_schema: default_schema.to_string(),
supports_session_database_switch: sqlserver_supports_session_database_switch(engine_edition),
})
}
pub fn completion_context_sql() -> &'static str {
SQLSERVER_COMPLETION_CONTEXT_SQL
}
pub fn completion_context_from_query_result(result: QueryResult) -> Result<SqlServerCompletionContext, String> {
let row = result.rows.first().ok_or_else(|| "SQL Server completion context query returned no rows".to_string())?;
let default_schema = row.first().and_then(serde_json::Value::as_str);
let engine_edition = row.get(1).and_then(|value| {
value
.as_i64()
.and_then(|value| i32::try_from(value).ok())
.or_else(|| value.as_str()?.trim().parse::<i32>().ok())
});
sqlserver_completion_context(default_schema, engine_edition)
}
#[derive(Debug, PartialEq, Eq)]
struct SqlServerEndpoint<'a> {
host: &'a str,
@ -1133,6 +1183,15 @@ pub async fn list_databases(client: &mut SqlServerClient) -> Result<Vec<Database
Ok(rows.iter().map(|row| DatabaseInfo { name: row.get::<&str, _>(0).unwrap_or("").to_string() }).collect())
}
pub async fn get_completion_context(client: &mut SqlServerClient) -> Result<SqlServerCompletionContext, String> {
let stream = client.query(SQLSERVER_COMPLETION_CONTEXT_SQL, &[]).await.map_err(|e| e.to_string())?;
let rows = stream.into_first_result().await.map_err(|e| e.to_string())?;
let row = rows.first().ok_or_else(|| "SQL Server completion context query returned no rows".to_string())?;
let default_schema = row.try_get::<&str, _>(0).map_err(|e| e.to_string())?;
let engine_edition = row.try_get::<i32, _>(1).map_err(|e| e.to_string())?;
sqlserver_completion_context(default_schema, engine_edition)
}
pub async fn test_connection(client: &mut SqlServerClient) -> Result<(), String> {
crate::db::with_connection_timeout("SQL Server", crate::db::connection_timeout(), async {
let stream = client.simple_query("SELECT 1").await.map_err(|e| e.to_string())?;
@ -2503,17 +2562,18 @@ fn first_sql_tokens(sql: &str, limit: usize) -> Vec<String> {
#[cfg(test)]
mod tests {
use super::{
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_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,
sqlserver_list_tables_sql, sqlserver_probe_explicit_alias, sqlserver_schema_name_predicate,
build_sqlserver_unsafe_type_query, capture_sqlserver_messages, completion_context_from_query_result,
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_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, sqlserver_list_tables_sql,
sqlserver_probe_explicit_alias, sqlserver_schema_name_predicate, sqlserver_supports_session_database_switch,
sqlserver_table_comment_sql, sqlserver_triggers_sql, sqlserver_visible_object_predicate,
strip_dbx_sqlserver_row_number_column, SqlServerDescribedColumn, SqlServerProbeOutputNameOverride,
SqlServerResultSet, SQLSERVER_RESULT_TYPE_PROBE_SQL,
SqlServerResultSet, SQLSERVER_COMPLETION_CONTEXT_SQL, SQLSERVER_RESULT_TYPE_PROBE_SQL,
};
use crate::types::{
CompletionAssistantMatchMode, CompletionAssistantObjectKind, CompletionAssistantRequest, QueryResult,
@ -2965,6 +3025,46 @@ mod tests {
assert!(sqlserver_indexes_sql("", "orders").contains("OBJECT_ID(QUOTENAME(N'orders'))"));
}
#[test]
fn sqlserver_completion_context_uses_the_database_user_default_schema() {
assert!(SQLSERVER_COMPLETION_CONTEXT_SQL.contains("SCHEMA_NAME()"));
assert!(SQLSERVER_COMPLETION_CONTEXT_SQL.contains("sys.schemas"));
assert!(SQLSERVER_COMPLETION_CONTEXT_SQL.contains("N'dbo'"));
assert!(SQLSERVER_COMPLETION_CONTEXT_SQL.contains("EngineEdition"));
}
#[test]
fn sqlserver_completion_context_disables_use_for_azure_database_endpoints() {
assert!(!sqlserver_supports_session_database_switch(5));
assert!(!sqlserver_supports_session_database_switch(6));
assert!(!sqlserver_supports_session_database_switch(11));
assert!(!sqlserver_supports_session_database_switch(12));
assert!(!sqlserver_supports_session_database_switch(99));
assert!(sqlserver_supports_session_database_switch(3));
assert!(sqlserver_supports_session_database_switch(8));
assert!(sqlserver_supports_session_database_switch(9));
}
#[test]
fn sqlserver_completion_context_parses_agent_query_results() {
let context = completion_context_from_query_result(QueryResult {
columns: vec!["default_schema".to_string(), "engine_edition".to_string()],
column_types: vec![],
column_sortables: vec![],
rows: vec![vec![serde_json::json!("app_user"), serde_json::json!("8")]],
affected_rows: 0,
execution_time_ms: 0,
truncated: false,
session_id: None,
has_more: false,
elasticsearch_raw_body: None,
})
.unwrap();
assert_eq!(context.default_schema, "app_user");
assert!(context.supports_session_database_switch);
}
#[test]
fn sqlserver_metadata_preserves_explicit_schema_matching() {
let predicate = sqlserver_schema_name_predicate("Sales'Ops", "s.name");

View File

@ -236,6 +236,55 @@ pub async fn list_sqlserver_linked_servers_core(
Ok(vec![])
}
pub async fn get_sqlserver_completion_context_core(
state: &AppState,
connection_id: &str,
database: &str,
) -> Result<db::sqlserver::SqlServerCompletionContext, String> {
retry_metadata_connection(state, connection_id, Some(database), || async {
let pool_key = state.get_or_create_pool(connection_id, Some(database)).await?;
let db_config = connection_config(state, connection_id).await;
let connections = state.connections.read().await;
if let Some(PoolKind::ExternalDriver { config, session, .. }) = connections.get(&pool_key) {
let config = config.clone();
let session = session.clone();
drop(connections);
let result: db::QueryResult = session
.invoke_with_timeout(
"executeQuery",
serde_json::json!({
"connection": config.as_ref(),
"database": database,
"sql": db::sqlserver::completion_context_sql(),
"maxRows": 1
}),
agent_metadata_timeout(Some(config.as_ref())),
)
.await?;
return db::sqlserver::completion_context_from_query_result(result);
}
try_sqlserver!(connections, &pool_key, get_completion_context);
if let Some(client) = extract_pool!(&connections, &pool_key, Agent) {
drop(connections);
let mut client = client.lock().await;
let result = client
.execute_query_with_timeout::<db::QueryResult>(
agent_execute_query_params(
db::sqlserver::completion_context_sql(),
if database.is_empty() { None } else { Some(database) },
None,
QueryExecutionOptions { max_rows: Some(1), ..Default::default() },
),
agent_metadata_timeout(db_config.as_ref()),
)
.await?;
return db::sqlserver::completion_context_from_query_result(result);
}
Err("SQL Server completion context requires a SQL Server connection".to_string())
})
.await
}
pub async fn list_sqlserver_linked_server_catalogs_core(
state: &AppState,
connection_id: &str,

View File

@ -344,6 +344,7 @@ async fn main() {
// Schema
.route("/schema/databases", get(routes::schema::list_databases))
.route("/schema/database-storage", post(routes::schema::list_database_storage))
.route("/schema/sqlserver/completion-context", get(routes::schema::get_sqlserver_completion_context))
.route("/schema/doris/catalogs", get(routes::schema::list_doris_catalogs))
.route("/schema/doris/catalog-databases", get(routes::schema::list_doris_catalog_databases))
.route("/schema/sqlserver/linked-servers", get(routes::schema::list_sqlserver_linked_servers))

View File

@ -52,6 +52,17 @@ pub async fn list_database_storage(
Ok(Json(result))
}
pub async fn get_sqlserver_completion_context(
State(state): State<Arc<WebState>>,
Query(q): Query<SchemaQuery>,
) -> Result<Json<dbx_core::db::sqlserver::SqlServerCompletionContext>, AppError> {
let database = q.database.as_deref().unwrap_or("");
let result = dbx_core::schema::get_sqlserver_completion_context_core(&state.app, &q.connection_id, database)
.await
.map_err(AppError::from)?;
Ok(Json(result))
}
/// Resolve a non-internal catalog for dispatch to the Doris multi-catalog path.
async fn external_doris_catalog(state: &Arc<WebState>, connection_id: &str, catalog: Option<&str>) -> Option<String> {
dbx_core::schema::resolve_external_doris_catalog(&state.app, connection_id, catalog).await

View File

@ -28,6 +28,15 @@ pub async fn list_database_storage(
dbx_core::schema::list_database_storage_core(&state, &connection_id, &databases).await
}
#[tauri::command]
pub async fn get_sqlserver_completion_context(
state: State<'_, Arc<AppState>>,
connection_id: String,
database: String,
) -> Result<db::sqlserver::SqlServerCompletionContext, String> {
dbx_core::schema::get_sqlserver_completion_context_core(&state, &connection_id, &database).await
}
#[tauri::command]
pub async fn list_doris_catalogs(
state: State<'_, Arc<AppState>>,

View File

@ -1633,6 +1633,7 @@ pub fn run() {
commands::plugins::uninstall_jdbc_plugin,
commands::schema::list_databases,
commands::schema::list_database_storage,
commands::schema::get_sqlserver_completion_context,
commands::schema::list_doris_catalogs,
commands::schema::list_doris_catalog_databases,
commands::schema::list_sqlserver_linked_servers,