diff --git a/app/admin/middleware/db.go b/app/admin/middleware/db.go index 79a32da9..0ff1c3bf 100644 --- a/app/admin/middleware/db.go +++ b/app/admin/middleware/db.go @@ -1,5 +1,58 @@ package middleware -import "go-admin/common/middleware" +import ( + "database/sql" + "errors" + "go-admin/common/config" + "go-admin/common/global" + "go-admin/tools" + "gorm.io/gorm/schema" + + "gorm.io/driver/mysql" + "gorm.io/driver/postgres" + "gorm.io/gorm" + + "go-admin/common/middleware" +) var WithContextDb = middleware.WithContextDb + +func getGormFromDb(driver string, db *sql.DB, config *gorm.Config) (*gorm.DB, error) { + switch driver { + case "mysql": + return gorm.Open(mysql.New(mysql.Config{Conn: db}), config) + case "postgres": + return gorm.Open(postgres.New(postgres.Config{Conn: db}), config) + default: + return nil, errors.New("not support this db driver") + } +} + +func GetGormFromConfig(cfg config.Conf) map[string]*gorm.DB { + gormDB := make(map[string]*gorm.DB) + if cfg.GetSaas() { + var err error + for k, v := range cfg.GetDbs() { + gormDB[k], err = getGormFromDb(v.Driver, v.DB, &gorm.Config{ + NamingStrategy: schema.NamingStrategy{ + SingularTable: true, + }, + }) + if err != nil { + global.Logger.Fatal(tools.Red(k+" connect error :"), err) + } + } + return gormDB + } + c := cfg.GetDb() + db, err := getGormFromDb(c.Driver, c.DB, &gorm.Config{ + NamingStrategy: schema.NamingStrategy{ + SingularTable: true, + }, + }) + if err != nil { + global.Logger.Fatal(tools.Red(c.Driver+" connect error :"), err) + } + gormDB["*"] = db + return gormDB +} diff --git a/app/admin/router/initrouter.go b/app/admin/router/initrouter.go index e6411dab..48176d79 100644 --- a/app/admin/router/initrouter.go +++ b/app/admin/router/initrouter.go @@ -8,20 +8,13 @@ import ( _ "go-admin/pkg/jwtauth" "go-admin/tools" config2 "go-admin/tools/config" - "gorm.io/gorm" ) -func InitRouter() *gin.Engine { - var r *gin.Engine - if global.GinEngine == nil { - r = gin.New() - } else { - r = global.GinEngine - } +func InitRouter(r *gin.Engine) *gin.Engine { if config2.SslConfig.Enable { r.Use(handler.TlsHandler()) } - r.Use(middleware.WithContextDb(map[string]*gorm.DB{"*": global.Eloquent})) + r.Use(middleware.WithContextDb(middleware.GetGormFromConfig(global.Cfg))) middleware.InitMiddleware(r) // the jwt middleware authMiddleware, err := middleware.AuthInit() diff --git a/cmd/api/server.go b/cmd/api/server.go index 25f12e00..1f8d1aac 100644 --- a/cmd/api/server.go +++ b/cmd/api/server.go @@ -69,13 +69,26 @@ func run() error { if viper.GetString("settings.application.mode") == string(tools.ModeProd) { gin.SetMode(gin.ReleaseMode) } + engine := global.Cfg.GetEngine() + if engine == nil { + engine = gin.New() + } - r := router.InitRouter() + var r *gin.Engine + switch engine.(type) { + case *gin.Engine: + r = engine.(*gin.Engine) + default: + panic("not support this engine") + } + + r = router.InitRouter(r) //defer global.Eloquent.Close() + global.Cfg.SetEngine(r) srv := &http.Server{ Addr: config.ApplicationConfig.Host + ":" + config.ApplicationConfig.Port, - Handler: r, + Handler: global.Cfg.GetEngine(), } go func() { jobs.InitJob() diff --git a/common/config/config.go b/common/config/config.go new file mode 100644 index 00000000..4ac318dd --- /dev/null +++ b/common/config/config.go @@ -0,0 +1,79 @@ +package config + +import ( + "database/sql" + "net/http" + + "go-admin/logger" +) + +type Config struct { + saas bool + dbs map[string]*DBConfig + db *DBConfig + engine http.Handler +} + +type DBConfig struct { + Driver string + DB *sql.DB +} + +// SetDbs 设置对应key的db +func (c *Config) SetDbs(key string, db *DBConfig) { + c.dbs[key] = db +} + +// GetDbs 获取所有map里的db数据 +func (c *Config) GetDbs() map[string]*DBConfig { + return c.dbs +} + +// GetDbByKey 根据key获取db +func (c *Config) GetDbByKey(key string) *DBConfig { + return c.dbs[key] +} + +// SetDb 设置单个db +func (c *Config) SetDb(db *DBConfig) { + c.db = db +} + +// GetDb 获取单个db +func (c *Config) GetDb() *DBConfig { + return c.db +} + +// SetEngine 设置路由引擎 +func (c *Config) SetEngine(engine http.Handler) { + c.engine = engine +} + +// GetEngine 获取路由引擎 +func (c *Config) GetEngine() http.Handler { + return c.engine +} + +// SetLogger 设置日志组件 +func (c *Config) SetLogger(l logger.Logger) { + logger.DefaultLogger = l +} + +// GetLogger 获取日志组件 +func (c *Config) GetLogger() logger.Logger { + return logger.DefaultLogger +} + +// SetSaas 设置是否是saas应用 +func (c *Config) SetSaas(saas bool) { + c.saas = saas +} + +// GetSaas 获取是否是saas应用 +func (c *Config) GetSaas() bool { + return c.saas +} + +func DefaultConfig() *Config { + return &Config{} +} diff --git a/common/config/type.go b/common/config/type.go new file mode 100644 index 00000000..8fbed122 --- /dev/null +++ b/common/config/type.go @@ -0,0 +1,28 @@ +package config + +import ( + "net/http" + + "go-admin/logger" +) + +type Conf interface { + //多db设置,⚠️SetDbs不允许并发,可以根据自己的业务,例如app分库、host分库 + SetDbs(key string, db *DBConfig) + GetDbs() map[string]*DBConfig + GetDbByKey(key string) *DBConfig + GetSaas() bool + SetSaas(bool) + + //单库业务实现这两个接口 + SetDb(db *DBConfig) + GetDb() *DBConfig + + //使用的路由 + SetEngine(engine http.Handler) + GetEngine() http.Handler + + //使用go-admin定义的logger,参考来源go-micro + SetLogger(logger logger.Logger) + GetLogger() logger.Logger +} diff --git a/common/database/mysql-drive.go b/common/database/mysql-drive.go index a053321f..fa83e47a 100644 --- a/common/database/mysql-drive.go +++ b/common/database/mysql-drive.go @@ -1,28 +1,37 @@ package database import ( - "gorm.io/driver/mysql" - "gorm.io/gorm" - "gorm.io/gorm/logger" - "gorm.io/gorm/schema" + "database/sql" "log" "os" "time" + "gorm.io/driver/mysql" + "gorm.io/gorm" + "gorm.io/gorm/logger" + "gorm.io/gorm/schema" + + "go-admin/common/config" "go-admin/common/global" "go-admin/tools" - "go-admin/tools/config" + toolsConfig "go-admin/tools/config" ) type Mysql struct { } func (e *Mysql) Setup() { - var err error - global.Source = e.GetConnect() global.Logger.Info(tools.Green(global.Source)) - global.Eloquent, err = e.Open(e.GetConnect(), &gorm.Config{ + db, err := sql.Open("mysql", global.Source) + if err != nil { + global.Logger.Fatal(tools.Red(e.GetDriver()+" connect error :"), err) + } + global.Cfg.SetDb(&config.DBConfig{ + Driver: "mysql", + DB: db, + }) + global.Eloquent, err = e.Open(db, &gorm.Config{ NamingStrategy: schema.NamingStrategy{ SingularTable: true, }, @@ -37,7 +46,7 @@ func (e *Mysql) Setup() { global.Logger.Fatal(tools.Red(" database error :"), global.Eloquent.Error) } - if config.LoggerConfig.EnabledDB { + if toolsConfig.LoggerConfig.EnabledDB { global.Eloquent.Logger = logger.New(log.New(os.Stdout, "\r\n", log.LstdFlags), logger.Config{ SlowThreshold: time.Second, Colorful: true, @@ -47,15 +56,15 @@ func (e *Mysql) Setup() { } // 打开数据库连接 -func (e *Mysql) Open(conn string, cfg *gorm.Config) (db *gorm.DB, err error) { - return gorm.Open(mysql.Open(conn), cfg) +func (e *Mysql) Open(db *sql.DB, cfg *gorm.Config) (*gorm.DB, error) { + return gorm.Open(mysql.New(mysql.Config{Conn: db}), cfg) } // 获取数据库连接 func (e *Mysql) GetConnect() string { - return config.DatabaseConfig.Source + return toolsConfig.DatabaseConfig.Source } func (e *Mysql) GetDriver() string { - return config.DatabaseConfig.Driver + return toolsConfig.DatabaseConfig.Driver } diff --git a/common/database/pgsql-driver.go b/common/database/pgsql-driver.go index 7cbce140..e3980e13 100644 --- a/common/database/pgsql-driver.go +++ b/common/database/pgsql-driver.go @@ -1,6 +1,7 @@ package database import ( + "database/sql" "log" "os" "time" @@ -10,8 +11,10 @@ import ( "gorm.io/gorm/logger" "gorm.io/gorm/schema" + "go-admin/common/config" "go-admin/common/global" - "go-admin/tools/config" + "go-admin/tools" + toolsConfig "go-admin/tools/config" ) type PgSql struct { @@ -22,7 +25,15 @@ func (e *PgSql) Setup() { global.Source = e.GetConnect() log.Println(global.Source) - global.Eloquent, err = e.Open(e.GetDriver(), &gorm.Config{ + db, err := sql.Open("postgresql", global.Source) + if err != nil { + global.Logger.Fatal(tools.Red(e.GetDriver()+" connect error :"), err) + } + global.Cfg.SetDb(&config.DBConfig{ + Driver: "mysql", + DB: db, + }) + global.Eloquent, err = e.Open(db, &gorm.Config{ NamingStrategy: schema.NamingStrategy{ SingularTable: true, }, @@ -37,7 +48,7 @@ func (e *PgSql) Setup() { log.Fatalf("database error %v", global.Eloquent.Error) } - if config.LoggerConfig.EnabledDB { + if toolsConfig.LoggerConfig.EnabledDB { global.Eloquent.Logger = logger.New(log.New(os.Stdout, "\r\n", log.LstdFlags), logger.Config{ SlowThreshold: time.Second, Colorful: true, @@ -47,15 +58,14 @@ func (e *PgSql) Setup() { } // 打开数据库连接 -func (*PgSql) Open(conn string, cfg *gorm.Config) (db *gorm.DB, err error) { - eloquent, err := gorm.Open(postgres.Open(conn), cfg) - return eloquent, err +func (e *PgSql) Open(db *sql.DB, cfg *gorm.Config) (*gorm.DB, error) { + return gorm.Open(postgres.New(postgres.Config{Conn: db}), cfg) } func (e *PgSql) GetConnect() string { - return config.DatabaseConfig.Source + return toolsConfig.DatabaseConfig.Source } func (e *PgSql) GetDriver() string { - return config.DatabaseConfig.Driver + return toolsConfig.DatabaseConfig.Driver } diff --git a/common/global/adm.go b/common/global/adm.go index bdb1c597..9ef4e8d7 100644 --- a/common/global/adm.go +++ b/common/global/adm.go @@ -5,9 +5,17 @@ import ( "github.com/gin-gonic/gin" "github.com/gogf/gf/os/glog" "github.com/robfig/cron/v3" + "go-admin/common/config" "gorm.io/gorm" ) +const ( + // go-admin Version Info + Version = "1.2.0" +) + +var Cfg config.Conf = config.DefaultConfig() + var GinEngine *gin.Engine var CasbinEnforcer *casbin.SyncedEnforcer var Eloquent *gorm.DB @@ -20,13 +28,6 @@ var ( DBName string ) -// go-admin Version Info -var Version string - -func init() { - Version = "1.2.0" -} - var ( Logger *glog.Logger JobLogger *glog.Logger