fix(mysql): switch tab database label after USE statements

This commit is contained in:
zipg
2026-09-22 11:48:36 +08:00
committed by GitHub
parent 713ce06aea
commit 8b3d7fc9e6
7 changed files with 174 additions and 16 deletions
@@ -12,6 +12,8 @@ import {
sqlServerLeadingUseScript,
sqlServerUseCompletionDatabaseNames,
sqlServerUseDatabaseFromStatement,
switchesDatabaseWithUseStatement,
useDatabaseFromStatement,
} from "@/lib/sql/sqlCompletionLookupTarget";
describe("sqlCompletionLookupTarget", () => {
@@ -290,6 +292,35 @@ describe("sqlCompletionLookupTarget", () => {
expect(sqlServerLeadingUseScript("USE FooDB; INSERT INTO audit_log VALUES (1); SELECT * FROM Users;")).toBeUndefined();
});
it("parses the USE target of MySQL-family dialects in all three identifier forms", () => {
expect(useDatabaseFromStatement("USE hd_ods", "mysql")).toBe("hd_ods");
expect(useDatabaseFromStatement("use hd_ods;", "mysql")).toBe("hd_ods");
expect(useDatabaseFromStatement(" USE `hd-ods`; ", "doris")).toBe("hd-ods");
expect(useDatabaseFromStatement('USE "报告库";', "starrocks")).toBe("报告库");
expect(useDatabaseFromStatement("-- switch\n/* c */ USE goldendb_db;", "goldendb")).toBe("goldendb_db");
expect(useDatabaseFromStatement("USE `a``b`;", "gbase")).toBe("a`b");
});
it("does not turn other statements or dialects into a database switch", () => {
expect(useDatabaseFromStatement("SELECT 1 FROM dual;", "mysql")).toBeUndefined();
expect(useDatabaseFromStatement("USE hd_ods; SELECT 1;", "mysql")).toBeUndefined();
expect(useDatabaseFromStatement("USE hd_ods;", "postgres")).toBeUndefined();
expect(useDatabaseFromStatement("USE other_ods;", "oracle")).toBeUndefined();
// SQL Server keeps its own bracket-aware parser.
expect(useDatabaseFromStatement("USE [Bar]]DB];", "sqlserver")).toBe("Bar]DB");
expect(useDatabaseFromStatement("USE hd_ods;", undefined)).toBeUndefined();
});
it("reports which dialects treat USE as a database switch", () => {
for (const databaseType of ["mysql", "doris", "starrocks", "goldendb", "gbase"] as const) {
expect(switchesDatabaseWithUseStatement(databaseType)).toBe(true);
}
for (const databaseType of ["postgres", "oracle", "clickhouse", "hive"] as const) {
expect(switchesDatabaseWithUseStatement(databaseType)).toBe(false);
}
expect(switchesDatabaseWithUseStatement(undefined)).toBe(false);
});
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;
@@ -54,6 +54,32 @@ export function sqlServerUseDatabaseFromStatement(statement: string): string | u
return match[3];
}
/**
* 方言里 `USE <db>` 会把会话切到另一个库,DBX 的标签库名应当跟着走(#9941)。
*
* 只列出已经确认过这一语义的方言:MySQL 及其 wire-protocol 家族。SQL Server 由
* `sqlServerUseDatabaseFromStatement` 单独处理(它还接受 `[db]` 括号标识符)。
*/
const USE_DATABASE_SWITCH_DIALECTS: ReadonlySet<DatabaseType> = new Set<DatabaseType>(["mysql", "doris", "starrocks", "goldendb", "gbase"]);
export function switchesDatabaseWithUseStatement(databaseType: DatabaseType | null | undefined): boolean {
return !!databaseType && USE_DATABASE_SWITCH_DIALECTS.has(databaseType);
}
/**
* 解析单条语句里「成功切换当前库」的 `USE <db>`,返回目标库名(反引号、双引号或裸
* 标识符)。不是 USE 语句、或该方言的 USE 不改库时返回 undefined。
*/
export function useDatabaseFromStatement(statement: string, databaseType?: DatabaseType): string | undefined {
if (databaseType === "sqlserver") return sqlServerUseDatabaseFromStatement(statement);
if (!switchesDatabaseWithUseStatement(databaseType)) return 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;
@@ -0,0 +1,96 @@
import { createPinia, setActivePinia } from "pinia";
import { beforeEach, describe, expect, it, vi } from "vitest";
import type { QueryResult } from "@/types/database";
const mocks = vi.hoisted(() => ({
executeMulti: vi.fn(),
executeQuery: vi.fn(),
closeClientConnectionSession: vi.fn(),
saveOpenTabsState: vi.fn(),
preparePaginationPlan: vi.fn(),
getConfig: vi.fn(),
ensureConnected: vi.fn(),
}));
vi.mock("@/lib/backend/api", () => ({
executeMulti: mocks.executeMulti,
executeQuery: mocks.executeQuery,
closeClientConnectionSession: mocks.closeClientConnectionSession,
saveOpenTabsState: mocks.saveOpenTabsState,
prepareQueryPaginationExecutionPlan: mocks.preparePaginationPlan,
}));
vi.mock("@/stores/connectionStore", () => ({
useConnectionStore: () => ({
getConfig: mocks.getConfig,
ensureConnected: mocks.ensureConnected,
recordConnectionLostError: vi.fn(),
}),
}));
vi.mock("@/stores/settingsStore", () => ({
useSettingsStore: () => ({
editorSettings: { pageSize: 100, openTabsRestoreMode: "all", confirmUnsavedSqlClose: false },
}),
}));
const okResult: QueryResult = { columns: [], rows: [], affected_rows: 0, execution_time_ms: 1 };
function installLocalStorage() {
const data = new Map<string, string>();
vi.stubGlobal("localStorage", {
getItem: vi.fn((key: string) => data.get(key) ?? null),
setItem: vi.fn((key: string, value: string) => data.set(key, value)),
removeItem: vi.fn((key: string) => data.delete(key)),
});
}
describe("queryStore USE database statement", () => {
beforeEach(() => {
vi.clearAllMocks();
vi.unstubAllGlobals();
installLocalStorage();
setActivePinia(createPinia());
mocks.getConfig.mockReturnValue({ id: "mysql-1", name: "MySQL", db_type: "mysql", query_timeout_secs: 45 });
mocks.executeMulti.mockResolvedValue([okResult]);
mocks.executeQuery.mockResolvedValue(okResult);
mocks.closeClientConnectionSession.mockResolvedValue(true);
mocks.saveOpenTabsState.mockResolvedValue(undefined);
mocks.ensureConnected.mockResolvedValue(undefined);
mocks.preparePaginationPlan.mockImplementation(async (options: { sql: string }) => ({ sqlToExecute: options.sql }));
});
it("follows a successful USE statement in the tab's database", async () => {
const { useQueryStore } = await import("@/stores/queryStore");
const store = useQueryStore();
const tabId = store.createTab("mysql-1", "dbx", "Query", "query");
await store.executeTabSql(tabId, "USE dbx_dup;");
expect(mocks.executeMulti).toHaveBeenCalledTimes(1);
expect(store.tabs.find((tab) => tab.id === tabId)?.database).toBe("dbx_dup");
await vi.waitFor(() => expect(mocks.closeClientConnectionSession).toHaveBeenCalled());
});
it("keeps the tab on its database when the USE statement fails", async () => {
mocks.executeMulti.mockResolvedValue([{ ...okResult, execution_error: true }]);
const { useQueryStore } = await import("@/stores/queryStore");
const store = useQueryStore();
const tabId = store.createTab("mysql-1", "dbx", "Query", "query");
await store.executeTabSql(tabId, "USE missing_db;");
expect(store.tabs.find((tab) => tab.id === tabId)?.database).toBe("dbx");
});
it("does not follow USE for dialects where it is not a database switch", async () => {
mocks.getConfig.mockReturnValue({ id: "pg-1", name: "PostgreSQL", db_type: "postgres", query_timeout_secs: 45 });
const { useQueryStore } = await import("@/stores/queryStore");
const store = useQueryStore();
const tabId = store.createTab("pg-1", "public_db", "Query", "query");
await store.executeTabSql(tabId, "USE other_db;");
expect(store.tabs.find((tab) => tab.id === tabId)?.database).toBe("public_db");
});
});
+21 -16
View File
@@ -67,7 +67,7 @@ import { normalizeResultPageSize } from "@/lib/dataGrid/paginationPageSize";
import { agentProtocolQueryResultMaxRows, capQueryResultTotal, effectiveQueryResultMaxRows, limitQueryPagination, queryResultLimitReached } from "@/lib/dataGrid/queryResultRowLimit";
import { elasticsearchRestRequestRanges, executableStatementRanges, splitSqlStatementRanges, sqlStatementParameterOptionsForCompatibility, stripMysqlClientDisplayCommand } from "@/lib/sql/sqlStatementRanges";
import type { SqlParameterOptions } from "@/lib/sql/sqlParameters";
import { replaceSqlServerLeadingUseQuery, sqlServerLeadingUseScript, sqlServerUseDatabaseFromStatement } from "@/lib/sql/sqlCompletionLookupTarget";
import { replaceSqlServerLeadingUseQuery, sqlServerLeadingUseScript, switchesDatabaseWithUseStatement, useDatabaseFromStatement } from "@/lib/sql/sqlCompletionLookupTarget";
import { classifySqlRisk } from "@/lib/sql/sqlRisk";
import { externalSqlFileDisplayTitles, normalizeExternalSqlPath } from "@/lib/sql/sqlFileOpen";
import { clearDataGridPendingSnapshot, clearDataGridPendingSnapshotsForTab } from "@/composables/useDataGridEditor";
@@ -369,16 +369,7 @@ function preservedResultIndex(results: QueryResult[], currentIndex: number | und
return currentIndex;
}
function annotateQueryResultSources(
results: QueryResult[],
sql: string,
database: string | undefined,
databaseType?: DatabaseType,
sourceOffset?: number,
parameterOptions?: SqlParameterOptions,
executedSql?: string,
sourceDocumentSql?: string,
): { results: QueryResult[]; sqlServerUseDatabase?: string } {
function annotateQueryResultSources(results: QueryResult[], sql: string, database: string | undefined, databaseType?: DatabaseType, sourceOffset?: number, parameterOptions?: SqlParameterOptions, executedSql?: string, sourceDocumentSql?: string): { results: QueryResult[]; useDatabase?: string } {
const statements = splitSqlStatementRanges(sql, databaseType, parameterOptions);
// The backend positions errors against the SQL it actually received. When the
// sent SQL was rewritten (pagination wrapper, injected hidden keys…), record
@@ -389,7 +380,7 @@ function annotateQueryResultSources(
const documentStatements = sourceDocumentSql && sourceOffset !== undefined ? splitSqlStatementRanges(sourceDocumentSql, databaseType, parameterOptions) : [];
let statementIndex = 0;
let sourceDatabase = database;
let sqlServerUseDatabase: string | undefined;
let useDatabase: string | undefined;
for (const result of results) {
const explicitIndex = Number.isInteger(result.statement_index) && result.statement_index! >= 0 ? result.statement_index : undefined;
const sourceIndex = explicitIndex ?? statementIndex;
@@ -420,13 +411,13 @@ function annotateQueryResultSources(
const preamble = documentStatement ? sourceDocumentSql!.slice(documentStatement.hitFrom, documentStatement.from) : sql.slice(statement.hitFrom, statement.from);
const customName = queryResultNameFromPreamble(preamble, { databaseType });
if (customName) result.sourceLabel = customName;
const successfulUseDatabase = databaseType === "sqlserver" && result.execution_error !== true ? sqlServerUseDatabaseFromStatement(statement.sql) : undefined;
const successfulUseDatabase = result.execution_error !== true ? useDatabaseFromStatement(statement.sql, databaseType) : undefined;
if (successfulUseDatabase) {
sourceDatabase = successfulUseDatabase;
sqlServerUseDatabase = successfulUseDatabase;
useDatabase = successfulUseDatabase;
}
}
return { results, sqlServerUseDatabase };
return { results, useDatabase };
}
/**
@@ -7557,7 +7548,11 @@ export const useQueryStore = defineStore("query", () => {
}
const successfulOracleSchemaChanges = usesOracleStickyTransactionState(effectiveDbType) ? results.filter((result) => result.execution_error !== true && isOracleCurrentSchemaStatement(result.sourceStatement)).length : 0;
const successfulSapHanaSchemaChanges = effectiveDbType === "saphana" ? results.filter((result) => result.execution_error !== true && isSapHanaSetSchemaStatement(result.sourceStatement)).length : 0;
const sqlServerUseDatabase = effectiveDbType === "sqlserver" ? annotatedResults.sqlServerUseDatabase : undefined;
const sqlServerUseDatabase = effectiveDbType === "sqlserver" ? annotatedResults.useDatabase : undefined;
// MySQL 家族(含 Doris/StarRocks)的 `USE db` 同样会切走会话的当前库,标签库名
// 要跟着走,否则工具栏、标签标题和侧栏仍指向旧库(#9941)。SQL Server 走上面的
// 分支,它有额外的事务与 reset 语义。
const mysqlUseDatabase = switchesDatabaseWithUseStatement(effectiveDbType) ? annotatedResults.useDatabase : undefined;
if (hiddenPrimaryKeys.length > 0 && results.length === 1) {
const hiddenIndexes = hiddenResultColumnIndexes(results[0]!.columns, hiddenPrimaryKeys);
if (hiddenIndexes.length > 0) results[0]!.hidden_column_indexes = hiddenIndexes;
@@ -7601,6 +7596,16 @@ export const useQueryStore = defineStore("query", () => {
current.database = sqlServerUseDatabase;
current.schema = undefined;
}
if (mysqlUseDatabase && !usesExternalExecutionTarget && current.database !== mysqlUseDatabase) {
// 切库后旧库的池(池按「连接 + 库」分桶)不再被这个标签复用,旧会话却已经在
// server 端停在新库上;不关掉它,用户切回旧库时会被重新用上,出现「标签写着 A、
// 实际在 B」的错配。标签上挂着的显式事务在切库后同样不可达,一并收掉(与 SQL
// Server 分支一致)。
rollbackTabTransaction(current);
void closeClientConnectionSession(current);
current.database = mysqlUseDatabase;
current.schema = undefined;
}
const activeGroupIndex = current.activeResultIndex;
const activeGroupResults = current.results;
const shouldAppendResult = !!options?.appendResult && !!current.result;
Binary file not shown.

After

Width:  |  Height:  |  Size: 234 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 237 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 235 KiB