From 6393721d7a3307d0f52089faad36ddbfa3653692 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=80=97=E5=AD=90?= Date: Sat, 7 Dec 2024 04:09:50 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E6=95=B0=E6=8D=AE=E5=BA=93=E9=93=BE?= =?UTF-8?q?=E6=8E=A5=E5=86=85=E5=AD=98=E6=B3=84=E6=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/apps/mysql/service.go | 1 + internal/data/backup.go | 4 ++++ internal/data/database.go | 7 +++++++ internal/data/database_server.go | 11 ++++++++--- internal/data/database_user.go | 16 ++++++++++++++-- pkg/db/mysql.go | 3 +++ pkg/db/postgres.go | 7 +++++-- 7 files changed, 42 insertions(+), 7 deletions(-) diff --git a/internal/apps/mysql/service.go b/internal/apps/mysql/service.go index 18eb40c05..0b3003c22 100644 --- a/internal/apps/mysql/service.go +++ b/internal/apps/mysql/service.go @@ -199,6 +199,7 @@ func (s *Service) SetRootPassword(w http.ResponseWriter, r *http.Request) { return } } else { + defer mysql.Close() if err = mysql.UserPassword("root", req.Password, "localhost"); err != nil { service.Error(w, http.StatusInternalServerError, "%v", err) return diff --git a/internal/data/backup.go b/internal/data/backup.go index 74e614e07..0c34401db 100644 --- a/internal/data/backup.go +++ b/internal/data/backup.go @@ -244,6 +244,7 @@ func (r *backupRepo) createMySQL(to string, name string) error { if err != nil { return err } + defer mysql.Close() if exist, _ := mysql.DatabaseExists(name); !exist { return fmt.Errorf("数据库不存在:%s", name) } @@ -287,6 +288,7 @@ func (r *backupRepo) createPostgres(to string, name string) error { if err != nil { return err } + defer postgres.Close() if exist, _ := postgres.DatabaseExist(name); !exist { return fmt.Errorf("数据库不存在:%s", name) } @@ -400,6 +402,7 @@ func (r *backupRepo) restoreMySQL(backup, target string) error { if err != nil { return err } + defer mysql.Close() if exist, _ := mysql.DatabaseExists(target); !exist { return fmt.Errorf("数据库不存在:%s", target) } @@ -435,6 +438,7 @@ func (r *backupRepo) restorePostgres(backup, target string) error { if err != nil { return err } + defer postgres.Close() if exist, _ := postgres.DatabaseExist(target); !exist { return fmt.Errorf("数据库不存在:%s", target) } diff --git a/internal/data/database.go b/internal/data/database.go index 3bf286945..74ba8a117 100644 --- a/internal/data/database.go +++ b/internal/data/database.go @@ -42,6 +42,7 @@ func (r databaseRepo) List(page, limit uint) ([]*biz.Database, int64, error) { }) } } + _ = mysql.Close() } case biz.DatabaseTypePostgresql: postgres, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port) @@ -58,6 +59,7 @@ func (r databaseRepo) List(page, limit uint) ([]*biz.Database, int64, error) { }) } } + _ = postgres.Close() } } } @@ -77,6 +79,7 @@ func (r databaseRepo) Create(req *request.DatabaseCreate) error { if err != nil { return err } + defer mysql.Close() if req.CreateUser { if err = NewDatabaseUserRepo().Create(&request.DatabaseUserCreate{ ServerID: req.ServerID, @@ -100,6 +103,7 @@ func (r databaseRepo) Create(req *request.DatabaseCreate) error { if err != nil { return err } + defer postgres.Close() if req.CreateUser { if err = NewDatabaseUserRepo().Create(&request.DatabaseUserCreate{ ServerID: req.ServerID, @@ -138,12 +142,14 @@ func (r databaseRepo) Delete(serverID uint, name string) error { if err != nil { return err } + defer mysql.Close() return mysql.DatabaseDrop(name) case biz.DatabaseTypePostgresql: postgres, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port) if err != nil { return err } + defer postgres.Close() return postgres.DatabaseDrop(name) } @@ -164,6 +170,7 @@ func (r databaseRepo) Comment(req *request.DatabaseComment) error { if err != nil { return err } + defer postgres.Close() return postgres.DatabaseComment(req.Name, req.Comment) } diff --git a/internal/data/database_server.go b/internal/data/database_server.go index fdfbbcd6f..cd6f563d4 100644 --- a/internal/data/database_server.go +++ b/internal/data/database_server.go @@ -130,6 +130,7 @@ func (r databaseServerRepo) Sync(id uint) error { if err != nil { return err } + defer mysql.Close() allUsers, err := mysql.Users() if err != nil { return err @@ -154,6 +155,7 @@ func (r databaseServerRepo) Sync(id uint) error { if err != nil { return err } + defer postgres.Close() allUsers, err := postgres.Users() if err != nil { return err @@ -179,20 +181,23 @@ func (r databaseServerRepo) Sync(id uint) error { func (r databaseServerRepo) checkServer(server *biz.DatabaseServer) bool { switch server.Type { case biz.DatabaseTypeMysql: - _, err := db.NewMySQL(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) + mysql, err := db.NewMySQL(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err == nil { + _ = mysql.Close() server.Status = biz.DatabaseServerStatusValid return true } case biz.DatabaseTypePostgresql: - _, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port) + postgres, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port) if err == nil { + _ = postgres.Close() server.Status = biz.DatabaseServerStatusValid return true } case biz.DatabaseTypeRedis: - _, err := db.NewRedis(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) + redis, err := db.NewRedis(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err == nil { + _ = redis.Close() server.Status = biz.DatabaseServerStatusValid return true } diff --git a/internal/data/database_user.go b/internal/data/database_user.go index 57c59c106..333f1d338 100644 --- a/internal/data/database_user.go +++ b/internal/data/database_user.go @@ -63,6 +63,7 @@ func (r databaseUserRepo) Create(req *request.DatabaseUserCreate) error { if err != nil { return err } + defer mysql.Close() if err = mysql.UserCreate(req.Username, req.Password, req.Host); err != nil { return err } @@ -81,6 +82,7 @@ func (r databaseUserRepo) Create(req *request.DatabaseUserCreate) error { if err != nil { return err } + defer postgres.Close() if err = postgres.UserCreate(req.Username, req.Password); err != nil { return err } @@ -122,6 +124,7 @@ func (r databaseUserRepo) Update(req *request.DatabaseUserUpdate) error { if err != nil { return err } + defer mysql.Close() if req.Password != "" { if err = mysql.UserPassword(user.Username, req.Password, user.Host); err != nil { return err @@ -137,6 +140,7 @@ func (r databaseUserRepo) Update(req *request.DatabaseUserUpdate) error { if err != nil { return err } + defer postgres.Close() if req.Password != "" { if err = postgres.UserPassword(user.Username, req.Password); err != nil { return err @@ -183,12 +187,14 @@ func (r databaseUserRepo) Delete(id uint) error { if err != nil { return err } + defer mysql.Close() _ = mysql.UserDrop(user.Username, user.Host) case biz.DatabaseTypePostgresql: postgres, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port) if err != nil { return err } + defer postgres.Close() _ = postgres.UserDrop(user.Username) } @@ -207,6 +213,7 @@ func (r databaseUserRepo) DeleteByNames(serverID uint, names []string) error { if err != nil { return err } + defer mysql.Close() users := make([]*biz.DatabaseUser, 0) if err = app.Orm.Where("server_id = ? AND username IN ?", serverID, names).Find(&users).Error; err != nil { return err @@ -226,6 +233,7 @@ func (r databaseUserRepo) DeleteByNames(serverID uint, names []string) error { if err != nil { return err } + defer postgres.Close() for name := range slices.Values(names) { _ = postgres.UserDrop(name) } @@ -246,10 +254,12 @@ func (r databaseUserRepo) fillUser(user *biz.DatabaseUser) { case biz.DatabaseTypeMysql: mysql, err := db.NewMySQL(server.Username, server.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)) if err == nil { + defer mysql.Close() privileges, _ := mysql.UserPrivileges(user.Username, user.Host) user.Privileges = privileges } - if _, err := db.NewMySQL(user.Username, user.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)); err == nil { + if mysql2, err := db.NewMySQL(user.Username, user.Password, fmt.Sprintf("%s:%d", server.Host, server.Port)); err == nil { + _ = mysql2.Close() user.Status = biz.DatabaseUserStatusValid } else { user.Status = biz.DatabaseUserStatusInvalid @@ -257,10 +267,12 @@ func (r databaseUserRepo) fillUser(user *biz.DatabaseUser) { case biz.DatabaseTypePostgresql: postgres, err := db.NewPostgres(server.Username, server.Password, server.Host, server.Port) if err == nil { + defer postgres.Close() privileges, _ := postgres.UserPrivileges(user.Username) user.Privileges = privileges } - if _, err := db.NewPostgres(user.Username, user.Password, server.Host, server.Port); err == nil { + if postgres2, err := db.NewPostgres(user.Username, user.Password, server.Host, server.Port); err == nil { + _ = postgres2.Close() user.Status = biz.DatabaseUserStatusValid } else { user.Status = biz.DatabaseUserStatusInvalid diff --git a/pkg/db/mysql.go b/pkg/db/mysql.go index b1ab10e89..9d369d8ff 100644 --- a/pkg/db/mysql.go +++ b/pkg/db/mysql.go @@ -3,6 +3,7 @@ package db import ( "database/sql" "fmt" + "net/url" "regexp" "slices" @@ -19,6 +20,8 @@ type MySQL struct { } func NewMySQL(username, password, address string, typ ...string) (*MySQL, error) { + username = url.QueryEscape(username) + password = url.QueryEscape(password) dsn := fmt.Sprintf("%s:%s@tcp(%s)/", username, password, address) if len(typ) > 0 && typ[0] == "unix" { dsn = fmt.Sprintf("%s:%s@unix(%s)/", username, password, address) diff --git a/pkg/db/postgres.go b/pkg/db/postgres.go index 426b48a02..1d93c4810 100644 --- a/pkg/db/postgres.go +++ b/pkg/db/postgres.go @@ -4,6 +4,7 @@ import ( "database/sql" "fmt" "slices" + "strings" _ "github.com/lib/pq" @@ -20,12 +21,14 @@ type Postgres struct { } func NewPostgres(username, password, address string, port uint) (*Postgres, error) { - dsn := fmt.Sprintf("host=%s port=%d user=%s password=%s dbname=postgres sslmode=disable", address, port, username, password) + username = strings.ReplaceAll(username, `'`, `\'`) + password = strings.ReplaceAll(password, `'`, `\'`) + dsn := fmt.Sprintf(`host=%s port=%d user='%s' password='%s' dbname=postgres sslmode=disable`, address, port, username, password) if password == "" { if username == "" { username = "postgres" } - dsn = fmt.Sprintf("host=%s port=%d user=%s dbname=postgres sslmode=disable", address, port, username) + dsn = fmt.Sprintf(`host=%s port=%d user='%s' dbname=postgres sslmode=disable`, address, port, username) } db, err := sql.Open("postgres", dsn) if err != nil {