From 486a73ef09f6db2680f7c42ba7d2b59172f672dd Mon Sep 17 00:00:00 2001 From: guoyongchang Date: Fri, 31 Jul 2026 21:22:31 +0800 Subject: [PATCH] fix(sqlserver): honor USE context in completions --- .../src/components/editor/QueryEditor.vue | 308 +++++++++++++----- .../queryEditorDatabaseNameCompletion.spec.ts | 19 ++ .../sql/sqlCompletionLookupTarget.spec.ts | 245 +++++++++++++- apps/desktop/src/lib/backend/api.ts | 1 + apps/desktop/src/lib/backend/http.ts | 5 + apps/desktop/src/lib/backend/tauri.ts | 5 + .../src/lib/sql/sqlCompletionLookupTarget.ts | 190 ++++++++++- .../connectionStore.completion.spec.ts | 97 +++++- apps/desktop/src/stores/connectionStore.ts | 24 +- apps/desktop/src/types/database.ts | 5 + crates/dbx-core/src/db/sqlserver.rs | 118 ++++++- crates/dbx-core/src/schema.rs | 49 +++ crates/dbx-web/src/main.rs | 1 + crates/dbx-web/src/routes/schema.rs | 11 + src-tauri/src/commands/schema.rs | 9 + src-tauri/src/lib.rs | 1 + 16 files changed, 974 insertions(+), 114 deletions(-) diff --git a/apps/desktop/src/components/editor/QueryEditor.vue b/apps/desktop/src/components/editor/QueryEditor.vue index 912879a92..f548c22d5 100644 --- a/apps/desktop/src/components/editor/QueryEditor.vue +++ b/apps/desktop/src/components/editor/QueryEditor.vue @@ -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(); // Persistent column cache keyed by "schema.table" or "table" const cachedColumnsByTable = new Map(); const cachedInsertValueHintColumnsByTable = new Map(); @@ -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; + +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, databaseNames: string[]): string[] { - if (!props.connectionId || props.database == null || !mayCompleteDatabaseSchemaQualifier(completionContext)) return []; +function localCompletionSchemasForDatabaseDisambiguation(completionContext: ReturnType, 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, fullDoc: string, position: number) { +function buildLocalSqlCompletionResult(completionContext: ReturnType, 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(); 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 0) { columnsByTable.set(cacheKey, localColumns); @@ -2915,7 +3010,7 @@ function buildLocalSqlCompletionResult(completionContext: ReturnType, fullDoc: string, position: number) { +function scheduleCompletionMetadataRefresh(completionContext: ReturnType, 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 { const scopedTables = tables.map((table) => ({ ...table, database: table.database ?? tableLookupTarget.database })); cachedTables = mergeCompletionTables(cachedTables, scopedTables); @@ -2960,11 +3055,13 @@ function scheduleCompletionMetadataRefresh(completionContext: ReturnType {}); } 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 { - 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 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(remote: Promise, local: T): Promise((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 { +function listCompletionTablesWithLatencyBudget(connectionId: string, database: string, filter: string, limit: number, schema?: string, globalSearch = false, catalog = props.catalog, currentSchema = props.schema): Promise { 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): 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, scope: CompletionMetadataScope) { + return resolveSqlCompletionRoutineLookupTarget({ + currentDatabase: scope.database, + currentSchema: scope.schema, + supportsDatabaseSchemaQualifier: supportsDatabaseSchemaQualifierCompletion(), + completionContext, + }); } -async function listCompletionObjectsForContext(completionContext: ReturnType): Promise { +function routineCompletionScopeForContext(completionContext: ReturnType, 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, 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, scope: CompletionMetadataScope): Promise { 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, fullDoc: string, position: number) { +async function performAsyncCompletionWithResult(epoch: number, completionContext: ReturnType, 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(); 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) { diff --git a/apps/desktop/src/lib/__tests__/editor/queryEditorDatabaseNameCompletion.spec.ts b/apps/desktop/src/lib/__tests__/editor/queryEditorDatabaseNameCompletion.spec.ts index e407e0819..f587d5f84 100644 --- a/apps/desktop/src/lib/__tests__/editor/queryEditorDatabaseNameCompletion.spec.ts +++ b/apps/desktop/src/lib/__tests__/editor/queryEditorDatabaseNameCompletion.spec.ts @@ -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"); + }); }); diff --git a/apps/desktop/src/lib/__tests__/sql/sqlCompletionLookupTarget.spec.ts b/apps/desktop/src/lib/__tests__/sql/sqlCompletionLookupTarget.spec.ts index 453693437..40ca5ce9a 100644 --- a/apps/desktop/src/lib/__tests__/sql/sqlCompletionLookupTarget.spec.ts +++ b/apps/desktop/src/lib/__tests__/sql/sqlCompletionLookupTarget.spec.ts @@ -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_", + }); + }); }); diff --git a/apps/desktop/src/lib/backend/api.ts b/apps/desktop/src/lib/backend/api.ts index 95896650b..b07a2f7bc 100644 --- a/apps/desktop/src/lib/backend/api.ts +++ b/apps/desktop/src/lib/backend/api.ts @@ -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"); diff --git a/apps/desktop/src/lib/backend/http.ts b/apps/desktop/src/lib/backend/http.ts index 33136e220..21638726a 100644 --- a/apps/desktop/src/lib/backend/http.ts +++ b/apps/desktop/src/lib/backend/http.ts @@ -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 { + return get(`/api/schema/sqlserver/completion-context?${qs({ connection_id: connectionId, database })}`); +} + export async function listDorisCatalogs(connectionId: string): Promise { return get(`/api/schema/doris/catalogs?${qs({ connection_id: connectionId })}`); } diff --git a/apps/desktop/src/lib/backend/tauri.ts b/apps/desktop/src/lib/backend/tauri.ts index d49e2d9ae..0f74de832 100644 --- a/apps/desktop/src/lib/backend/tauri.ts +++ b/apps/desktop/src/lib/backend/tauri.ts @@ -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 { + return invoke("get_sqlserver_completion_context", { connectionId, database }); +} + export async function listDorisCatalogs(connectionId: string): Promise { return invoke("list_doris_catalogs", { connectionId }); } diff --git a/apps/desktop/src/lib/sql/sqlCompletionLookupTarget.ts b/apps/desktop/src/lib/sql/sqlCompletionLookupTarget.ts index 6b6fb5ed8..285a300f5 100644 --- a/apps/desktop/src/lib/sql/sqlCompletionLookupTarget.ts +++ b/apps/desktop/src/lib/sql/sqlCompletionLookupTarget.ts @@ -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 }): 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 }): 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, }; diff --git a/apps/desktop/src/stores/__tests__/connectionStore.completion.spec.ts b/apps/desktop/src/stores/__tests__/connectionStore.completion.spec.ts index a666f3792..dbe0b05fa 100644 --- a/apps/desktop/src/stores/__tests__/connectionStore.completion.spec.ts +++ b/apps/desktop/src/stores/__tests__/connectionStore.completion.spec.ts @@ -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(), deferred(), deferred(), deferred()]; let activeColumns = 0; diff --git a/apps/desktop/src/stores/connectionStore.ts b/apps/desktop/src/stores/connectionStore.ts index 5353bc383..43e68ddd7 100644 --- a/apps/desktop/src/stores/connectionStore.ts +++ b/apps/desktop/src/stores/connectionStore.ts @@ -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>({}); const completionForeignKeysCache = ref>({}); const completionDatabasesCache = ref>({}); + const sqlServerCompletionContextCache = ref>({}); const elasticsearchCompletionIndicesCache = ref>({}); const redisCompletionKeysCache = ref>({}); const mongoCompletionCollectionsCache = ref>({}); @@ -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 { 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 { + 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 { const cacheKey = `${connectionId}:${database}`; if (schemaListCache.value[cacheKey]) { @@ -6776,6 +6795,7 @@ export const useConnectionStore = defineStore("connection", () => { listCompletionForeignKeys, listCompletionSchemas, listCompletionDatabases, + getSqlServerCompletionContext, lookupLocalCompletionTables, lookupLocalCompletionObjects, lookupLocalCompletionColumns, diff --git a/apps/desktop/src/types/database.ts b/apps/desktop/src/types/database.ts index b62b5c6ba..568078edb 100644 --- a/apps/desktop/src/types/database.ts +++ b/apps/desktop/src/types/database.ts @@ -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; diff --git a/crates/dbx-core/src/db/sqlserver.rs b/crates/dbx-core/src/db/sqlserver.rs index 75732b151..1bb9fc63d 100644 --- a/crates/dbx-core/src/db/sqlserver.rs +++ b/crates/dbx-core/src/db/sqlserver.rs @@ -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, +) -> Result { + 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 { + 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::().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(0).unwrap_or("").to_string() }).collect()) } +pub async fn get_completion_context(client: &mut SqlServerClient) -> Result { + 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::(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 { #[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"); diff --git a/crates/dbx-core/src/schema.rs b/crates/dbx-core/src/schema.rs index 4db7cd137..b96ff8ea4 100644 --- a/crates/dbx-core/src/schema.rs +++ b/crates/dbx-core/src/schema.rs @@ -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 { + 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::( + 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, diff --git a/crates/dbx-web/src/main.rs b/crates/dbx-web/src/main.rs index 32204283f..dc276b350 100644 --- a/crates/dbx-web/src/main.rs +++ b/crates/dbx-web/src/main.rs @@ -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)) diff --git a/crates/dbx-web/src/routes/schema.rs b/crates/dbx-web/src/routes/schema.rs index ce789d2dc..0d3b752b5 100644 --- a/crates/dbx-web/src/routes/schema.rs +++ b/crates/dbx-web/src/routes/schema.rs @@ -52,6 +52,17 @@ pub async fn list_database_storage( Ok(Json(result)) } +pub async fn get_sqlserver_completion_context( + State(state): State>, + Query(q): Query, +) -> Result, 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, connection_id: &str, catalog: Option<&str>) -> Option { dbx_core::schema::resolve_external_doris_catalog(&state.app, connection_id, catalog).await diff --git a/src-tauri/src/commands/schema.rs b/src-tauri/src/commands/schema.rs index 1a3929562..f45c260f5 100644 --- a/src-tauri/src/commands/schema.rs +++ b/src-tauri/src/commands/schema.rs @@ -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>, + connection_id: String, + database: String, +) -> Result { + dbx_core::schema::get_sqlserver_completion_context_core(&state, &connection_id, &database).await +} + #[tauri::command] pub async fn list_doris_catalogs( state: State<'_, Arc>, diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 2ab7e9fb3..bb5fa3cf8 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -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,