commit 32d7e9ac0dbf0305c50b73269c1296dfb674fbc3 Author: zyj <18107291228@163.com> Date: Thu Jul 17 16:55:32 2025 +0800 init diff --git a/config.yaml b/config.yaml new file mode 100644 index 0000000..df75aae --- /dev/null +++ b/config.yaml @@ -0,0 +1,11 @@ +# 数据库配置 +database: + username: "root" + password: "root" + host: "localhost" + port: 3306 + database: "jz" + +# 当前监听的端口 +port: 6310 + diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..e754eba --- /dev/null +++ b/go.mod @@ -0,0 +1,9 @@ +module mysql-operate + +go 1.24.3 + +require ( + filippo.io/edwards25519 v1.1.0 // indirect + github.com/go-sql-driver/mysql v1.9.3 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..f54656a --- /dev/null +++ b/go.sum @@ -0,0 +1,9 @@ +filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= +filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= +github.com/JamesStewy/go-mysqldump v0.2.2 h1:tMtZDnIi2hz6H3Nna0TPhvWfBZlXe4i7vkjc5Vd8Gdo= +github.com/JamesStewy/go-mysqldump v0.2.2/go.mod h1:JuJhv4dTbe2OQpABlwqj0B+6E9VLjGLG1t4NJRTcB3w= +github.com/go-sql-driver/mysql v1.9.3 h1:U/N249h2WzJ3Ukj8SowVFjdtZKfu9vlLZxjPXV1aweo= +github.com/go-sql-driver/mysql v1.9.3/go.mod h1:qn46aNg1333BRMNU69Lq93t8du/dwxI64Gl8i5p1WMU= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/main.go b/main.go new file mode 100644 index 0000000..a070739 --- /dev/null +++ b/main.go @@ -0,0 +1,59 @@ +package main + +import ( + "database/sql" + "flag" + "fmt" + "log" + "mysql-operate/service" + "strconv" + + _ "github.com/go-sql-driver/mysql" +) + +var ( + start = flag.Bool("start", false, "运行") + tabName = flag.String("tabName", "", "表名") + sqlFile = flag.String("sqlFile", "", "导出的sql文件名称,默认为数据库名称") +) + +func main() { + flag.Parse() + if !*start { + log.Println(*start) + return + } + appConfig, _ := service.LoadConfig() + + var sqlFileName string = appConfig.Database.Database + if *tabName != "" { + sqlFileName = *tabName + } + if *sqlFile != "" { + sqlFileName = *sqlFile + } + + // 获取数据库配置 + + dsn := fmt.Sprintf("%s:%s@tcp(%s:%s)/%s?charset=utf8mb4", + appConfig.Database.Username, + appConfig.Database.Password, + appConfig.Database.Host, + strconv.Itoa(appConfig.Database.Port), + appConfig.Database.Database, + ) + // 连接数据库 + db, err := sql.Open("mysql", dsn) + if err != nil { + panic(fmt.Sprintf("连接失败: %v", err)) + } + defer db.Close() + var msqldump service.Msqldump + dumo := msqldump.New(db, "./", sqlFileName) + if *tabName != "" { + dumo.Tablelist = []string{*tabName} + } + + dumo.Run() + +} diff --git a/service/dump.go b/service/dump.go new file mode 100644 index 0000000..cbf5008 --- /dev/null +++ b/service/dump.go @@ -0,0 +1,278 @@ +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 +}