mirror of
https://github.com/t8y2/dbx.git
synced 2026-10-02 02:34:42 +08:00
fix(mysql): switch tab database label after USE statements
This commit is contained in:
@@ -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");
|
||||
});
|
||||
});
|
||||
@@ -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 |
Reference in New Issue
Block a user