package service import ( "database/sql" "errors" "log" "os" "path" "strings" "text/template" "time" "gopkg.in/yaml.v3" ) type AppConfig struct { Port int `yaml:"port"` Database DatabaseConfig `yaml:"database"` } type DatabaseConfig struct { Username string `yaml:"username"` Password string `yaml:"password"` Host string `yaml:"host"` Port int `yaml:"port"` Database string `yaml:"database"` } type Dump struct { db *sql.DB Tablelist []string ServerVersion string FilePath string Host string Database string Tables []*table CompleteTime string DumpVersion string } type table struct { Name string SQL string Values string } type Msqldump struct{} const tmpl = `-- Go SQL Dump {{ .DumpVersion }} -- -- Host: {{ .Host }} Database: {{ .Database }} -- ------------------------------------------------------ -- Server version {{ .ServerVersion }} /*!40101 SET @OLD_CHARACTER_SET_CLIENT=@@CHARACTER_SET_CLIENT */; /*!40101 SET @OLD_CHARACTER_SET_RESULTS=@@CHARACTER_SET_RESULTS */; /*!40101 SET @OLD_COLLATION_CONNECTION=@@COLLATION_CONNECTION */; /*!40101 SET NAMES utf8 */; /*!40103 SET @OLD_TIME_ZONE=@@TIME_ZONE */; /*!40103 SET TIME_ZONE='+00:00' */; /*!40014 SET @OLD_UNIQUE_CHECKS=@@UNIQUE_CHECKS, UNIQUE_CHECKS=0 */; /*!40014 SET @OLD_FOREIGN_KEY_CHECKS=@@FOREIGN_KEY_CHECKS, FOREIGN_KEY_CHECKS=0 */; /*!40101 SET @OLD_SQL_MODE=@@SQL_MODE, SQL_MODE='NO_AUTO_VALUE_ON_ZERO' */; /*!40111 SET @OLD_SQL_NOTES=@@SQL_NOTES, SQL_NOTES=0 */; {{range .Tables}} -- -- Table structure for table {{ .Name }} -- DROP TABLE IF EXISTS {{ .Name }}; /*!40101 SET @saved_cs_client = @@character_set_client */; /*!40101 SET character_set_client = utf8mb4 */; {{ .SQL }}; /*!40101 SET character_set_client = @saved_cs_client */; -- -- Dumping data for table {{ .Name }} -- LOCK TABLES {{ .Name }} WRITE; /*!40000 ALTER TABLE {{ .Name }} DISABLE KEYS */; {{ if .Values }} INSERT INTO {{ .Name }} VALUES {{ .Values }}; {{ end }} /*!40000 ALTER TABLE {{ .Name }} ENABLE KEYS */; UNLOCK TABLES; {{ end }} /*!40103 SET TIME_ZONE=@OLD_TIME_ZONE */; -- Dump completed on {{ .CompleteTime }} ` func (*Msqldump) New(db *sql.DB, dir, fileName string) *Dump { appConfig, _ := LoadConfig() return &Dump{ db: db, Tablelist: []string{}, FilePath: path.Join(dir, fileName+".sql"), Host: appConfig.Database.Host, Database: appConfig.Database.Database, DumpVersion: "1.1.1", } } func (d *Dump) Run() error { // 获取当前版本 if serverVersion, err := getServerVersion(d.db); err != nil { d.ServerVersion = serverVersion } sqlFile, err := os.OpenFile(d.FilePath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0644) if err != nil { log.Println("sql文件打开失败") return err } defer sqlFile.Close() tables, _ := getTables(d.db, d.Tablelist) for _, name := range tables { if t, err := createTable(d.db, name); err == nil { d.Tables = append(d.Tables, t) } } for _, v := range d.Tables { v.Name = "`" + v.Name + "`" } d.CompleteTime = time.Now().String() t, err := template.New("mysqldump").Parse(tmpl) if err != nil { log.Println(err) } if err = t.Execute(sqlFile, d); err != nil { return err } return nil } // 获取表格列表 func getTables(db *sql.DB, tablelist []string) ([]string, error) { tables := make([]string, 0) if len(tablelist) > 0 { for _, v := range tablelist { var exists bool query := `SELECT COUNT(*) > 0 FROM information_schema.tables WHERE table_schema = DATABASE() AND table_name = ?` err := db.QueryRow(query, v).Scan(&exists) if err != nil { return nil, err } if exists { tables = append(tables, v) } } return tables, nil } rows, err := db.Query("SHOW TABLES") if err != nil { return tables, err } defer rows.Close() // Read result for rows.Next() { var table sql.NullString if err := rows.Scan(&table); err != nil { return tables, err } tables = append(tables, table.String) } return tables, rows.Err() } func createTable(db *sql.DB, name string) (*table, error) { var err error t := &table{Name: name} if t.SQL, err = createTableSQL(db, name); err != nil { return nil, err } if t.Values, err = createTableValues(db, name); err != nil { return nil, err } return t, nil } func createTableSQL(db *sql.DB, name string) (string, error) { // Get table creation SQL var table_return sql.NullString var table_sql sql.NullString err := db.QueryRow("SHOW CREATE TABLE "+name).Scan(&table_return, &table_sql) if err != nil { return "", err } if table_return.String != name { return "", errors.New("Returned table is not the same as requested table") } return table_sql.String, nil } func createTableValues(db *sql.DB, name string) (string, error) { // Get Data rows, err := db.Query("SELECT * FROM " + name) if err != nil { return "", err } defer rows.Close() // Get columns columns, err := rows.Columns() if err != nil { return "", err } if len(columns) == 0 { return "", errors.New("No columns in table " + name + ".") } // Read data data_text := make([]string, 0) for rows.Next() { // Init temp data storage //ptrs := make([]interface{}, len(columns)) //var ptrs []interface {} = make([]*sql.NullString, len(columns)) data := make([]*sql.NullString, len(columns)) ptrs := make([]interface{}, len(columns)) for i, _ := range data { ptrs[i] = &data[i] } // Read data if err := rows.Scan(ptrs...); err != nil { return "", err } dataStrings := make([]string, len(columns)) for key, value := range data { if value != nil && value.Valid { dataStrings[key] = "'" + value.String + "'" } else { dataStrings[key] = "null" } } data_text = append(data_text, "("+strings.Join(dataStrings, ",")+")") } return strings.Join(data_text, ","), rows.Err() } // 获取数据库版本 func getServerVersion(db *sql.DB) (string, error) { var server_version sql.NullString if err := db.QueryRow("SELECT version()").Scan(&server_version); err != nil { return "", err } return server_version.String, nil } func LoadConfig() (*AppConfig, error) { // 读取文件内容 data, err := os.ReadFile("config.yaml") if err != nil { return nil, err } // 解析 YAML var cfg AppConfig if err := yaml.Unmarshal(data, &cfg); err != nil { return nil, err } return &cfg, nil }