diff --git a/internal/app/methods_file.go b/internal/app/methods_file.go index 751c8d57..fb8d9df9 100644 --- a/internal/app/methods_file.go +++ b/internal/app/methods_file.go @@ -1550,6 +1550,74 @@ func executeSQLFileStream(ctx context.Context, dbInst db.Database, reader io.Rea // ExecuteSQLFile 在后端流式读取并执行大 SQL 文件,通过事件推送进度。 // 前端通过 EventsOn("sqlfile:progress", ...) 监听进度。 +const sqlFileExecutionPreambleBytes = 64 * 1024 + +func readSQLFileExecutionPreamble(reader io.ReadSeeker) ([]byte, error) { + buffer := make([]byte, sqlFileExecutionPreambleBytes) + read, err := io.ReadFull(reader, buffer) + if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, io.ErrUnexpectedEOF) { + return nil, err + } + if _, err := reader.Seek(0, io.SeekStart); err != nil { + return nil, err + } + return buffer[:read], nil +} + +type goNaviMySQLDatabaseBackupPreamble struct { + databaseName string + includesCreateDatabase bool +} + +func parseGoNaviMySQLDatabaseBackupPreamble(preamble []byte) (goNaviMySQLDatabaseBackupPreamble, bool) { + text := strings.TrimPrefix(string(preamble), "\ufeff") + if !strings.HasPrefix(strings.TrimSpace(text), "-- GoNavi SQL Export") { + return goNaviMySQLDatabaseBackupPreamble{}, false + } + + databaseName := "" + for _, line := range strings.Split(text, "\n") { + trimmed := strings.TrimSpace(line) + if strings.HasPrefix(trimmed, "-- Database:") { + databaseName = strings.TrimSpace(strings.TrimPrefix(trimmed, "-- Database:")) + break + } + } + if databaseName == "" { + return goNaviMySQLDatabaseBackupPreamble{}, false + } + + quotedDatabase := quoteIdentByType("mysql", databaseName) + if !strings.Contains(text, "USE "+quotedDatabase+";") { + return goNaviMySQLDatabaseBackupPreamble{}, false + } + return goNaviMySQLDatabaseBackupPreamble{ + databaseName: databaseName, + includesCreateDatabase: strings.Contains(text, "CREATE DATABASE IF NOT EXISTS "+quotedDatabase+";"), + }, true +} + +func buildGoNaviMySQLDatabaseBackupBootstrapSQL(backup goNaviMySQLDatabaseBackupPreamble) string { + if backup.includesCreateDatabase || strings.TrimSpace(backup.databaseName) == "" { + return "" + } + return fmt.Sprintf("CREATE DATABASE IF NOT EXISTS %s", quoteIdentByType("mysql", backup.databaseName)) +} + +func resolveSQLFileExecutionRunConfig(config connection.ConnectionConfig, dbName string, preamble []byte) connection.ConnectionConfig { + runConfig := normalizeRunConfig(config, dbName) + if strings.EqualFold(strings.TrimSpace(runConfig.Type), "mysql") { + _, isGoNaviDatabaseBackup := parseGoNaviMySQLDatabaseBackupPreamble(preamble) + if !isGoNaviDatabaseBackup { + return runConfig + } + // A GoNavi database backup creates and selects its source database itself. + // Connect at server level so restoring into a deleted database can start. + runConfig.Database = "" + } + return runConfig +} + func (a *App) ExecuteSQLFile(config connection.ConnectionConfig, dbName string, filePath string, jobID string) (result connection.QueryResult) { auditSQL := "EXECUTE SQL FILE" auditStatementCount := 0 @@ -1567,20 +1635,29 @@ func (a *App) ExecuteSQLFile(config connection.ConnectionConfig, dbName string, logger.Warnf("ExecuteSQLFile 开始:file=%s db=%s jobID=%s", filePath, dbName, jobID) - // 获取数据库连接 - runConfig := normalizeRunConfig(config, dbName) - dbInst, err := a.getDatabase(runConfig) - if err != nil { - logger.Error(err, "ExecuteSQLFile 获取连接失败:%s", formatConnSummary(runConfig)) - return connection.QueryResult{Success: false, Message: err.Error()} - } - // 打开文件 f, err := os.Open(filePath) if err != nil { return connection.QueryResult{Success: false, Message: a.appText("file.backend.error.open_file_failed", map[string]any{"detail": err.Error()})} } defer f.Close() + preamble, err := readSQLFileExecutionPreamble(f) + if err != nil { + return connection.QueryResult{Success: false, Message: a.appText("file.backend.error.open_file_failed", map[string]any{"detail": err.Error()})} + } + backupPreamble := goNaviMySQLDatabaseBackupPreamble{} + isGoNaviMySQLDatabaseBackup := false + if strings.EqualFold(strings.TrimSpace(config.Type), "mysql") { + backupPreamble, isGoNaviMySQLDatabaseBackup = parseGoNaviMySQLDatabaseBackupPreamble(preamble) + } + + // GoNavi 的 MySQL 整库备份会在脚本中创建并 USE 源库,因此不能先连接到该库。 + runConfig := resolveSQLFileExecutionRunConfig(config, dbName, preamble) + dbInst, err := a.getDatabase(runConfig) + if err != nil { + logger.Error(err, "ExecuteSQLFile 获取连接失败:%s", formatConnSummary(runConfig)) + return connection.QueryResult{Success: false, Message: err.Error()} + } // 获取文件大小用于计算进度 var totalSize int64 @@ -1606,6 +1683,12 @@ func (a *App) ExecuteSQLFile(config connection.ConnectionConfig, dbName string, a.queryMu.Unlock() }() + if bootstrapSQL := buildGoNaviMySQLDatabaseBackupBootstrapSQL(backupPreamble); isGoNaviMySQLDatabaseBackup && bootstrapSQL != "" { + if _, err := execSQLFileStatement(ctx, dbInst, bootstrapSQL); err != nil { + return connection.QueryResult{Success: false, Message: err.Error()} + } + } + // 发送进度事件的辅助函数 emitProgress := func(status string, executed, failed, total int, bytesRead int64, currentSQL string, errMsg string) { percent := 0.0 @@ -2881,7 +2964,7 @@ func (a *App) exportDatabaseSQLToFile( w := bufio.NewWriterSize(f, 1024*1024) defer w.Flush() - if err := writeSQLHeader(w, runConfig, dbName); err != nil { + if err := writeSQLDatabaseBackupHeader(w, runConfig, dbName); err != nil { return connection.QueryResult{Success: false, Message: err.Error()} } for _, objectName := range objects { @@ -3206,6 +3289,14 @@ func quoteQualifiedIdentByType(dbType string, ident string) string { } func writeSQLHeader(w *bufio.Writer, config connection.ConnectionConfig, dbName string) error { + return writeSQLHeaderWithDatabaseBootstrap(w, config, dbName, false) +} + +func writeSQLDatabaseBackupHeader(w *bufio.Writer, config connection.ConnectionConfig, dbName string) error { + return writeSQLHeaderWithDatabaseBootstrap(w, config, dbName, true) +} + +func writeSQLHeaderWithDatabaseBootstrap(w *bufio.Writer, config connection.ConnectionConfig, dbName string, createDatabase bool) error { now := time.Now().Format("2006-01-02 15:04:05") if _, err := w.WriteString(fmt.Sprintf("-- GoNavi SQL Export\n-- Time: %s\n", now)); err != nil { return err @@ -3217,6 +3308,11 @@ func writeSQLHeader(w *bufio.Writer, config connection.ConnectionConfig, dbName } if strings.ToLower(strings.TrimSpace(config.Type)) == "mysql" && strings.TrimSpace(dbName) != "" { + if createDatabase { + if _, err := w.WriteString(fmt.Sprintf("CREATE DATABASE IF NOT EXISTS %s;\n\n", quoteIdentByType("mysql", dbName))); err != nil { + return err + } + } if _, err := w.WriteString(fmt.Sprintf("USE %s;\n\n", quoteIdentByType("mysql", dbName))); err != nil { return err } diff --git a/internal/app/methods_file_export_test.go b/internal/app/methods_file_export_test.go index 286d3905..e7c0827a 100644 --- a/internal/app/methods_file_export_test.go +++ b/internal/app/methods_file_export_test.go @@ -1343,6 +1343,28 @@ func TestDumpTableSQL_OracleBackupBatchesRowsIntoInsertAll(t *testing.T) { } } +func TestWriteSQLDatabaseBackupHeaderCreatesMySQLDatabaseBeforeSelectingIt(t *testing.T) { + var output bytes.Buffer + writer := bufio.NewWriter(&output) + + if err := writeSQLDatabaseBackupHeader(writer, connection.ConnectionConfig{Type: "mysql"}, "restore_target"); err != nil { + t.Fatalf("writeSQLDatabaseBackupHeader returned error: %v", err) + } + if err := writer.Flush(); err != nil { + t.Fatalf("flush header: %v", err) + } + + content := output.String() + createIndex := strings.Index(content, "CREATE DATABASE IF NOT EXISTS `restore_target`;") + useIndex := strings.Index(content, "USE `restore_target`;") + if createIndex < 0 { + t.Fatalf("database backup header must create the source database, content=%q", content) + } + if useIndex < 0 || createIndex > useIndex { + t.Fatalf("database backup header must create the database before USE, content=%q", content) + } +} + func TestFilterExportObjectsBySchema_PostgresQualifiedObjectsOnly(t *testing.T) { got := filterExportObjectsBySchema( connection.ConnectionConfig{Type: "postgres"}, diff --git a/internal/app/methods_file_sql_execution_test.go b/internal/app/methods_file_sql_execution_test.go index 2d7bd456..275d21b0 100644 --- a/internal/app/methods_file_sql_execution_test.go +++ b/internal/app/methods_file_sql_execution_test.go @@ -584,3 +584,109 @@ func TestStreamSQLFileKeepsOraclePackageSpecAndBodyTogether(t *testing.T) { t.Fatalf("unexpected third statement: %q", statements[2]) } } + +func TestResolveSQLFileExecutionRunConfigUsesServerConnectionForGoNaviMySQLDatabaseBackup(t *testing.T) { + preamble := strings.Join([]string{ + "-- GoNavi SQL Export", + "-- Time: 2026-07-17 00:00:00", + "-- Database: restore_target", + "", + "CREATE DATABASE IF NOT EXISTS `restore_target`;", + "", + "USE `restore_target`;", + }, "\n") + + got := resolveSQLFileExecutionRunConfig( + connection.ConnectionConfig{Type: "mysql", Database: "selected_target"}, + "selected_target", + []byte(preamble), + ) + if got.Database != "" { + t.Fatalf("GoNavi MySQL database backup must connect at server level before CREATE/USE, got database=%q", got.Database) + } +} + +func TestResolveSQLFileExecutionRunConfigKeepsSelectedDatabaseForRegularSQL(t *testing.T) { + got := resolveSQLFileExecutionRunConfig( + connection.ConnectionConfig{Type: "mysql", Database: "configured_default"}, + "selected_target", + []byte("CREATE TABLE demo(id INT);"), + ) + if got.Database != "selected_target" { + t.Fatalf("regular SQL must retain the selected database, got database=%q", got.Database) + } +} + +func TestResolveSQLFileExecutionRunConfigUsesServerConnectionForLegacyGoNaviMySQLDatabaseBackup(t *testing.T) { + preamble := strings.Join([]string{ + "-- GoNavi SQL Export", + "-- Time: 2026-07-11 00:00:00", + "-- Database: legacy_restore_target", + "", + "USE `legacy_restore_target`;", + }, "\n") + + got := resolveSQLFileExecutionRunConfig( + connection.ConnectionConfig{Type: "mysql", Database: "selected_target"}, + "selected_target", + []byte(preamble), + ) + if got.Database != "" { + t.Fatalf("legacy GoNavi MySQL database backup must connect at server level before USE, got database=%q", got.Database) + } +} + +func TestBuildGoNaviMySQLDatabaseBackupBootstrapSQLOnlyForLegacyBackup(t *testing.T) { + legacy := goNaviMySQLDatabaseBackupPreamble{databaseName: "legacy_restore_target"} + if got := buildGoNaviMySQLDatabaseBackupBootstrapSQL(legacy); got != "CREATE DATABASE IF NOT EXISTS `legacy_restore_target`" { + t.Fatalf("unexpected legacy bootstrap SQL: %q", got) + } + + current := goNaviMySQLDatabaseBackupPreamble{ + databaseName: "current_restore_target", + includesCreateDatabase: true, + } + if got := buildGoNaviMySQLDatabaseBackupBootstrapSQL(current); got != "" { + t.Fatalf("backup that already creates its database must not be bootstrapped again, got %q", got) + } +} + +func TestExecuteSQLFileStreamRunsGoNaviMySQLDatabaseBackupHeader(t *testing.T) { + fakeDB := &fakeSQLFileBatchDB{} + input := strings.Join([]string{ + "-- GoNavi SQL Export", + "-- Database: restore_target", + "CREATE DATABASE IF NOT EXISTS `restore_target`;", + "USE `restore_target`;", + "SET FOREIGN_KEY_CHECKS=0;", + "CREATE TABLE users(id INT PRIMARY KEY);", + "INSERT INTO users(id) VALUES (1);", + "SET FOREIGN_KEY_CHECKS=1;", + }, "\n") + + result, err := executeSQLFileStream(context.Background(), fakeDB, strings.NewReader(input), sqlFileExecutionOptions{ + DBType: "mysql", + BatchMaxStatements: 100, + BatchMaxBytes: 1024, + }, nil) + if err != nil { + t.Fatalf("executeSQLFileStream returned error: %v", err) + } + if result.Executed != 6 || result.Failed != 0 { + t.Fatalf("expected complete database backup header and statements to execute, got %#v", result) + } + joinedExec := strings.Join(fakeDB.execQueries, "\n") + for _, expected := range []string{ + "CREATE DATABASE IF NOT EXISTS `restore_target`", + "USE `restore_target`", + "CREATE TABLE users(id INT PRIMARY KEY)", + "SET FOREIGN_KEY_CHECKS=1", + } { + if !strings.Contains(joinedExec, expected) { + t.Fatalf("expected backup statement %q to execute, queries=%#v", expected, fakeDB.execQueries) + } + } + if len(fakeDB.batchQueries) != 1 || !strings.Contains(fakeDB.batchQueries[0], "INSERT INTO users(id) VALUES (1)") { + t.Fatalf("expected INSERT data to be batched after schema restore, batches=%#v", fakeDB.batchQueries) + } +}