This commit is contained in:
zyj
2025-07-17 16:55:32 +08:00
commit 32d7e9ac0d
5 changed files with 366 additions and 0 deletions

278
service/dump.go Normal file
View File

@@ -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
}