From deeb82048bb0de9a9863cc4db32204ea8795f3b6 Mon Sep 17 00:00:00 2001 From: linwenxiang <991154416@qq.com> Date: Sun, 21 Feb 2021 14:06:13 +0800 Subject: [PATCH] =?UTF-8?q?=E6=95=B0=E6=8D=AE=E5=BA=93=E8=BF=9E=E6=8E=A5?= =?UTF-8?q?=E5=88=9D=E5=A7=8B=E5=8C=96=E4=BC=98=E5=8C=96=20=E6=9D=83?= =?UTF-8?q?=E9=99=90=E9=83=A8=E5=88=86=E4=BF=AE=E6=94=B9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- app/admin/apis/system/role.go | 167 ----------------- app/admin/apis/system/sys_role/sys_role.go | 124 +++++++++++-- app/admin/middleware/db.go | 42 ----- app/admin/middleware/db_other.go | 23 --- app/admin/middleware/db_sqlite3.go | 28 --- app/admin/middleware/init.go | 4 +- app/admin/middleware/permission.go | 4 +- app/admin/models/roledept.go | 47 ----- app/admin/models/system/role_dept.go | 11 ++ app/admin/models/system/role_menu.go | 172 ++++++++++++++++++ app/admin/router/initrouter.go | 1 - app/admin/router/sys_role.go | 1 + app/admin/router/sysrouter.go | 20 +- app/admin/service/dto/sys_role.go | 3 + app/admin/service/sys_role.go | 139 +++++++++++++- app/admin/service/sys_role_menu.go | 113 ++++++++++++ cmd/api/server.go | 24 +-- cmd/config/server.go | 4 +- .../migration/version/1599190683659_tables.go | 2 +- cmd/migrate/server.go | 2 +- common/config/config.go | 60 +++--- common/config/type.go | 15 +- common/database/initialize.go | 66 +++++-- common/database/initialize_sqlite3.go | 21 --- common/database/interface.go | 11 -- common/database/mysql_drive.go | 77 -------- common/database/open.go | 14 ++ common/database/open_sqlite3.go | 16 ++ common/database/pgsql_driver.go | 75 -------- common/database/sqlite3_driver.go | 77 -------- common/dto/search.go | 5 +- common/global/adm.go | 4 +- common/global/casbin.go | 12 +- common/middleware/db.go | 15 +- examples/run.go | 5 +- tools/config/config.go | 38 ++-- tools/config/database.go | 5 +- 37 files changed, 726 insertions(+), 721 deletions(-) delete mode 100644 app/admin/apis/system/role.go delete mode 100644 app/admin/middleware/db.go delete mode 100644 app/admin/middleware/db_other.go delete mode 100644 app/admin/middleware/db_sqlite3.go delete mode 100644 app/admin/models/roledept.go create mode 100644 app/admin/models/system/role_dept.go create mode 100644 app/admin/models/system/role_menu.go create mode 100644 app/admin/service/sys_role_menu.go delete mode 100644 common/database/initialize_sqlite3.go delete mode 100644 common/database/interface.go delete mode 100644 common/database/mysql_drive.go create mode 100644 common/database/open.go create mode 100644 common/database/open_sqlite3.go delete mode 100644 common/database/pgsql_driver.go delete mode 100644 common/database/sqlite3_driver.go diff --git a/app/admin/apis/system/role.go b/app/admin/apis/system/role.go deleted file mode 100644 index b2355647..00000000 --- a/app/admin/apis/system/role.go +++ /dev/null @@ -1,167 +0,0 @@ -package system - -import ( - "github.com/gin-gonic/gin" - "go-admin/app/admin/service/dto" - - "go-admin/app/admin/models" - "go-admin/common/global" - "go-admin/tools" - "go-admin/tools/app" -) - -// @Summary 角色列表数据 -// @Description Get JSON -// @Tags 角色/Role -// @Param roleName query string false "roleName" -// @Param status query string false "status" -// @Param roleKey query string false "roleKey" -// @Param pageSize query int false "页条数" -// @Param pageIndex query int false "页码" -// @Success 200 {object} app.Response "{"code": 200, "data": [...]}" -// @Router /api/v1/rolelist [get] -// @Security Bearer -func GetRoleList(c *gin.Context) { - var data models.SysRole - var err error - var pageSize = 10 - var pageIndex = 1 - - if size := c.Request.FormValue("pageSize"); size != "" { - pageSize, err = tools.StringToInt(size) - } - - if index := c.Request.FormValue("pageIndex"); index != "" { - pageIndex, err = tools.StringToInt(index) - } - - data.RoleKey = c.Request.FormValue("roleKey") - data.RoleName = c.Request.FormValue("roleName") - data.Status = c.Request.FormValue("status") - data.DataScope = tools.GetUserIdStr(c) - result, count, err := data.GetPage(pageSize, pageIndex) - tools.HasError(err, "", -1) - - app.PageOK(c, result, count, pageIndex, pageSize, "") -} - -// @Summary 获取Role数据 -// @Description 获取JSON -// @Tags 角色/Role -// @Param roleId path string false "roleId" -// @Success 200 {string} string "{"code": 200, "data": [...]}" -// @Success 200 {string} string "{"code": -1, "message": "抱歉未找到相关信息"}" -// @Router /api/v1/role [get] -// @Security Bearer -func GetRole(c *gin.Context) { - var Role models.SysRole - Role.RoleId, _ = tools.StringToInt(c.Param("roleId")) - result, err := Role.Get() - menuIds := make([]int, 0) - menuIds, err = Role.GetRoleMeunId() - tools.HasError(err, "抱歉未找到相关信息", -1) - result.MenuIds = menuIds - app.OK(c, result, "") - -} - -// @Summary 创建角色 -// @Description 获取JSON -// @Tags 角色/Role -// @Accept application/json -// @Product application/json -// @Param data body models.SysRole true "data" -// @Success 200 {string} string "{"code": 200, "message": "添加成功"}" -// @Success 200 {string} string "{"code": -1, "message": "添加失败"}" -// @Router /api/v1/role [post] -func InsertRole(c *gin.Context) { - var data models.SysRole - data.CreateBy = tools.GetUserIdStr(c) - err := c.Bind(&data) - tools.HasError(err, "数据解析失败", 500) - id, err := data.Insert() - tools.HasError(err, "", -1) - data.RoleId = id - - var t models.RoleMenu - if len(data.MenuIds) > 0 { - _, err = t.Insert(id, data.MenuIds) - tools.HasError(err, "", -1) - } - - _, err = global.LoadPolicy() - tools.HasError(err, "", -1) - - app.OK(c, data, "添加成功") -} - -// @Summary 修改用户角色 -// @Description 获取JSON -// @Tags 角色/Role -// @Accept application/json -// @Product application/json -// @Param data body models.SysRole true "body" -// @Success 200 {string} string "{"code": 200, "message": "修改成功"}" -// @Success 200 {string} string "{"code": -1, "message": "修改失败"}" -// @Router /api/v1/role [put] -func UpdateRole(c *gin.Context) { - var data models.SysRole - data.UpdateBy = tools.GetUserIdStr(c) - err := c.Bind(&data) - tools.HasError(err, "数据解析失败", -1) - result, err := data.Update(data.RoleId) - tools.HasError(err, "", -1) - var t models.RoleMenu - _, err = t.DeleteRoleMenu(data.RoleId) - tools.HasError(err, "修改失败(delete rm)", -1) - if len(data.MenuIds) > 0 { - _, err2 := t.Insert(data.RoleId, data.MenuIds) - tools.HasError(err2, "修改失败(insert)", -1) - } - - _, err = global.LoadPolicy() - tools.HasError(err, "", -1) - - app.OK(c, result, "修改成功") -} - -func UpdateRoleDataScope(c *gin.Context) { - var data models.SysRole - var req dto.RoleDataScopeReq - data.UpdateBy = tools.GetUserIdStr(c) - err := c.Bind(&req) - tools.HasError(err, "数据解析失败", -1) - data.RoleId = req.RoleId - data.DataScope = req.DataScope - result, err := data.Update(data.RoleId) - - var t models.SysRoleDept - _, err = t.DeleteRoleDept(data.RoleId) - tools.HasError(err, "添加失败1", -1) - if data.DataScope == "2" { - _, err2 := t.Insert(data.RoleId, data.DeptIds) - tools.HasError(err2, "添加失败2", -1) - } - app.OK(c, result, "修改成功") -} - -// @Summary 删除用户角色 -// @Description 删除数据 -// @Tags 角色/Role -// @Param roleId path int true "roleId" -// @Success 200 {string} string "{"code": 200, "message": "删除成功"}" -// @Success 200 {string} string "{"code": -1, "message": "删除失败"}" -// @Router /api/v1/role/{roleId} [delete] -func DeleteRole(c *gin.Context) { - var Role models.SysRole - Role.UpdateBy = tools.GetUserIdStr(c) - - IDS := tools.IdsStrToIdsIntGroup("roleId", c) - _, err := Role.BatchDelete(IDS) - tools.HasError(err, "删除失败", -1) - - _, err = global.LoadPolicy() - tools.HasError(err, "", -1) - - app.OK(c, "", "删除成功") -} diff --git a/app/admin/apis/system/sys_role/sys_role.go b/app/admin/apis/system/sys_role/sys_role.go index 2e535d74..c28af473 100644 --- a/app/admin/apis/system/sys_role/sys_role.go +++ b/app/admin/apis/system/sys_role/sys_role.go @@ -2,14 +2,13 @@ package sys_role import ( "github.com/gin-gonic/gin" - "go-admin/app/admin/models/system" "go-admin/app/admin/service" "go-admin/app/admin/service/dto" "go-admin/common/apis" + "go-admin/common/global" "go-admin/common/log" "go-admin/tools" - "net/http" ) @@ -17,6 +16,17 @@ type SysRole struct { apis.Api } +// @Summary 角色列表数据 +// @Description Get JSON +// @Tags 角色/Role +// @Param roleName query string false "roleName" +// @Param status query string false "status" +// @Param roleKey query string false "roleKey" +// @Param pageSize query int false "页条数" +// @Param pageIndex query int false "页码" +// @Success 200 {object} app.Response "{"code": 200, "data": [...]}" +// @Router /api/v1/role [get] +// @Security Bearer func (e *SysRole) GetSysRoleList(c *gin.Context) { msgID := tools.GenerateMsgIDFromContext(c) d := new(dto.SysRoleSearch) @@ -47,6 +57,14 @@ func (e *SysRole) GetSysRoleList(c *gin.Context) { e.PageOK(c, list, int(count), d.GetPageIndex(), d.GetPageSize(), "查询成功") } +// @Summary 获取Role数据 +// @Description 获取JSON +// @Tags 角色/Role +// @Param roleId path string false "roleId" +// @Success 200 {string} string "{"code": 200, "data": [...]}" +// @Success 200 {string} string "{"code": -1, "message": "抱歉未找到相关信息"}" +// @Router /api/v1/role/{id} [get] +// @Security Bearer func (e *SysRole) GetSysRole(c *gin.Context) { control := new(dto.SysRoleById) db, err := tools.GetOrm(c) @@ -76,6 +94,16 @@ func (e *SysRole) GetSysRole(c *gin.Context) { e.OK(c, object, "查看成功") } +// @Summary 创建角色 +// @Description 获取JSON +// @Tags 角色/Role +// @Accept application/json +// @Product application/json +// @Param data body models.SysRole true "data" +// @Success 200 {string} string "{"code": 200, "message": "添加成功"}" +// @Success 200 {string} string "{"code": -1, "message": "添加失败"}" +// @Router /api/v1/role [post] +// @Security Bearer func (e *SysRole) InsertSysRole(c *gin.Context) { control := new(dto.SysRoleControl) db, err := tools.GetOrm(c) @@ -97,21 +125,35 @@ func (e *SysRole) InsertSysRole(c *gin.Context) { return } // 设置创建人 - object.CreateBy= tools.GetUserId(c) + object.CreateBy = tools.GetUserId(c) - serviceSysRole := service.SysRole{} - serviceSysRole.Orm = db - serviceSysRole.MsgID = msgID - err = serviceSysRole.InsertSysRole(object) + s := service.SysRole{} + s.Orm = db + s.MsgID = msgID + err = s.InsertSysRole(object) if err != nil { log.Error(err) e.Error(c, http.StatusInternalServerError, err, "创建失败") return } - + _, err = global.LoadPolicy(c) + if err != nil { + e.Error(c, http.StatusInternalServerError, err, "") + return + } e.OK(c, object.GetId(), "创建成功") } +// @Summary 修改用户角色 +// @Description 获取JSON +// @Tags 角色/Role +// @Accept application/json +// @Product application/json +// @Param data body models.SysRole true "body" +// @Success 200 {string} string "{"code": 200, "message": "修改成功"}" +// @Success 200 {string} string "{"code": -1, "message": "修改失败"}" +// @Router /api/v1/role/{id} [put] +// @Security Bearer func (e *SysRole) UpdateSysRole(c *gin.Context) { control := new(dto.SysRoleControl) db, err := tools.GetOrm(c) @@ -134,17 +176,30 @@ func (e *SysRole) UpdateSysRole(c *gin.Context) { } object.UpdateBy = tools.GetUserId(c) - serviceSysRole := service.SysRole{} - serviceSysRole.Orm = db - serviceSysRole.MsgID = msgID - err = serviceSysRole.UpdateSysRole(object) + s := service.SysRole{} + s.Orm = db + s.MsgID = msgID + err = s.UpdateSysRole(object) if err != nil { log.Error(err) return } + _, err = global.LoadPolicy(c) + if err != nil { + e.Error(c, http.StatusInternalServerError, err, "") + return + } e.OK(c, object.GetId(), "更新成功") } +// @Summary 删除用户角色 +// @Description 删除数据 +// @Tags 角色/Role +// @Param roleId path int true "roleId" +// @Success 200 {string} string "{"code": 200, "message": "删除成功"}" +// @Success 200 {string} string "{"code": -1, "message": "删除失败"}" +// @Router /api/v1/role/{roleId} [delete] +// @Security Bearer func (e *SysRole) DeleteSysRole(c *gin.Context) { control := new(dto.SysRoleById) db, err := tools.GetOrm(c) @@ -162,14 +217,51 @@ func (e *SysRole) DeleteSysRole(c *gin.Context) { return } - serviceSysRole := service.SysRole{} - serviceSysRole.Orm = db - serviceSysRole.MsgID = msgID - err = serviceSysRole.RemoveSysRole(control) + s := service.SysRole{} + s.Orm = db + s.MsgID = msgID + err = s.RemoveSysRole(control) if err != nil { log.Error(err) return } + _, err = global.LoadPolicy(c) + if err != nil { + e.Error(c, http.StatusInternalServerError, err, "") + return + } e.OK(c, control.GetId(), "删除成功") } +func (e *SysRole) UpdateRoleDataScope(c *gin.Context) { + control := new(dto.RoleDataScopeReq) + db, err := tools.GetOrm(c) + if err != nil { + log.Error(err) + return + } + + msgID := tools.GenerateMsgIDFromContext(c) + //更新操作 + err = c.Bind(control) + if err != nil { + log.Errorf("msgID[%s] request bind error, %s", msgID, err.Error()) + e.Error(c, http.StatusUnprocessableEntity, err, "参数验证失败") + return + } + data := &system.SysRole{ + RoleId: control.RoleId, + DataScope: control.DataScope, + DeptIds: control.DeptIds, + } + data.UpdateBy = tools.GetUserId(c) + s := &service.SysRole{} + s.Orm = db + s.MsgID = msgID + err = s.UpdateDataScope(data) + if err != nil { + e.Error(c, http.StatusInternalServerError, err, "") + return + } + e.OK(c, nil, "操作成功") +} diff --git a/app/admin/middleware/db.go b/app/admin/middleware/db.go deleted file mode 100644 index 03f509ff..00000000 --- a/app/admin/middleware/db.go +++ /dev/null @@ -1,42 +0,0 @@ -package middleware - -import ( - "gorm.io/gorm" - "gorm.io/gorm/schema" - - "go-admin/common/config" - "go-admin/common/global" - "go-admin/common/middleware" - "go-admin/tools" -) - -var WithContextDb = middleware.WithContextDb - -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/middleware/db_other.go b/app/admin/middleware/db_other.go deleted file mode 100644 index d0eacbe9..00000000 --- a/app/admin/middleware/db_other.go +++ /dev/null @@ -1,23 +0,0 @@ -// +build !sqlite3 - -package middleware - -import ( - "database/sql" - "errors" - - "gorm.io/driver/mysql" - "gorm.io/driver/postgres" - "gorm.io/gorm" -) - -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") - } -} diff --git a/app/admin/middleware/db_sqlite3.go b/app/admin/middleware/db_sqlite3.go deleted file mode 100644 index f2e2bde0..00000000 --- a/app/admin/middleware/db_sqlite3.go +++ /dev/null @@ -1,28 +0,0 @@ -// +build sqlite3 - -package middleware - -import ( - "database/sql" - "errors" - - "gorm.io/driver/mysql" - "gorm.io/driver/postgres" - "gorm.io/driver/sqlite" - "gorm.io/gorm" - - "go-admin/common/global" -) - -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) - case "sqlite3": - return gorm.Open(sqlite.Open(global.Source), config) - default: - return nil, errors.New("not support this db driver") - } -} diff --git a/app/admin/middleware/init.go b/app/admin/middleware/init.go index 4b8be88e..7c15185d 100644 --- a/app/admin/middleware/init.go +++ b/app/admin/middleware/init.go @@ -17,5 +17,7 @@ func InitMiddleware(r *gin.Engine) { // Secure is a middleware function that appends security r.Use(Secure) // 链路追踪 - r.Use(middleware.Trace()) + //r.Use(middleware.Trace()) + // 数据库链接 + r.Use(middleware.WithContextDb) } diff --git a/app/admin/middleware/permission.go b/app/admin/middleware/permission.go index 8a5a138c..79187205 100644 --- a/app/admin/middleware/permission.go +++ b/app/admin/middleware/permission.go @@ -17,7 +17,7 @@ func AuthCheckRole() gin.HandlerFunc { return func(c *gin.Context) { data, _ := c.Get(jwtauth.JwtPayloadKey) v := data.(jwtauth.MapClaims) - e := global.CasbinEnforcer + e := global.Cfg.GetCasbinKey(c.Request.Host) var res bool var err error msgID := tools.GenerateMsgIDFromContext(c) @@ -38,7 +38,7 @@ func AuthCheckRole() gin.HandlerFunc { log.Infof("msgID[%s] isTrue: %v role: %s method: %s path: %s", msgID, res, v["rolekey"], c.Request.Method, c.Request.URL.Path) c.Next() } else { - log.Warnf("msgID[%s] isTrue: %v role: %s method: %s path: %s message: %s", msgID, res, v["rolekey"], c.Request.Method, c.Request.URL.Path,"当前request无权限,请管理员确认!") + log.Warnf("msgID[%s] isTrue: %v role: %s method: %s path: %s message: %s", msgID, res, v["rolekey"], c.Request.Method, c.Request.URL.Path, "当前request无权限,请管理员确认!") c.JSON(http.StatusOK, gin.H{ "code": 403, "msg": "对不起,您没有该接口访问权限,请联系管理员", diff --git a/app/admin/models/roledept.go b/app/admin/models/roledept.go deleted file mode 100644 index dd094e1c..00000000 --- a/app/admin/models/roledept.go +++ /dev/null @@ -1,47 +0,0 @@ -package models - -import ( - "fmt" - - orm "go-admin/common/global" -) - -//sys_role_dept -type SysRoleDept struct { - RoleId int `gorm:""` - DeptId int `gorm:""` -} - -func (SysRoleDept) TableName() string { - return "sys_role_dept" -} - -func (rm *SysRoleDept) Insert(roleId int, deptIds []int) (bool, error) { - //ORM不支持批量插入所以需要拼接 sql 串 - sql := "INSERT INTO `sys_role_dept` (`role_id`,`dept_id`) VALUES " - - for i := 0; i < len(deptIds); i++ { - if len(deptIds)-1 == i { - //最后一条数据 以分号结尾 - sql += fmt.Sprintf("(%d,%d);", roleId, deptIds[i]) - } else { - sql += fmt.Sprintf("(%d,%d),", roleId, deptIds[i]) - } - } - orm.Eloquent.Exec(sql) - - return true, nil -} - -func (rm *SysRoleDept) DeleteRoleDept(roleId int) (bool, error) { - if err := orm.Eloquent.Table("sys_role_dept").Where("role_id = ?", roleId).Delete(&rm).Error; err != nil { - return false, err - } - var role SysRole - if err := orm.Eloquent.Table("sys_role").Where("role_id = ?", roleId).First(&role).Error; err != nil { - return false, err - } - - return true, nil - -} diff --git a/app/admin/models/system/role_dept.go b/app/admin/models/system/role_dept.go new file mode 100644 index 00000000..526f42e7 --- /dev/null +++ b/app/admin/models/system/role_dept.go @@ -0,0 +1,11 @@ +package system + +//sys_role_dept +type SysRoleDept struct { + RoleId int `gorm:"size:11;primaryKey"` + DeptId int `gorm:"size:11;primaryKey"` +} + +func (SysRoleDept) TableName() string { + return "sys_role_dept" +} diff --git a/app/admin/models/system/role_menu.go b/app/admin/models/system/role_menu.go new file mode 100644 index 00000000..12f234fd --- /dev/null +++ b/app/admin/models/system/role_menu.go @@ -0,0 +1,172 @@ +package system + +import ( + "fmt" + "github.com/casbin/casbin/v2" + "gorm.io/gorm" + + "go-admin/app/admin/models" + "go-admin/tools" +) + +type RoleMenu struct { + RoleId int `gorm:""` + MenuId int `gorm:""` + RoleName string `gorm:"size:128)"` + CreateBy string `gorm:"size:128)"` + UpdateBy string `gorm:"size:128)"` +} + +func (RoleMenu) TableName() string { + return "sys_role_menu" +} + +type MenuPath struct { + Path string `json:"path"` +} + +func (rm *RoleMenu) Get(tx *gorm.DB) ([]RoleMenu, error) { + var r []RoleMenu + table := tx.Table("sys_role_menu") + if rm.RoleId != 0 { + table = table.Where("role_id = ?", rm.RoleId) + + } + if err := table.Find(&r).Error; err != nil { + return nil, err + } + return r, nil +} + +func (rm *RoleMenu) GetPermis(tx *gorm.DB) ([]string, error) { + var r []models.Menu + table := tx.Select("sys_menu.permission").Table("sys_menu").Joins("left join sys_role_menu on sys_menu.menu_id = sys_role_menu.menu_id") + + table = table.Where("role_id = ?", rm.RoleId) + + table = table.Where("sys_menu.menu_type in('F','C')") + if err := table.Find(&r).Error; err != nil { + return nil, err + } + var list []string + for i := 0; i < len(r); i++ { + list = append(list, r[i].Permission) + } + return list, nil +} + +func (rm *RoleMenu) GetIDS(tx *gorm.DB) ([]MenuPath, error) { + var r []MenuPath + table := tx.Select("sys_menu.path").Table("sys_role_menu") + table = table.Joins("left join sys_role on sys_role.role_id=sys_role_menu.role_id") + table = table.Joins("left join sys_menu on sys_menu.id=sys_role_menu.menu_id") + table = table.Where("sys_role.role_name = ? and sys_menu.type=1", rm.RoleName) + if err := table.Find(&r).Error; err != nil { + return nil, err + } + return r, nil +} + +func (rm *RoleMenu) DeleteRoleMenu(tx *gorm.DB, roleId int) error { + if err := tx.Table("sys_role_dept").Where("role_id = ?", roleId).Delete(&rm).Error; err != nil { + return err + } + if err := tx.Table("sys_role_menu").Where("role_id = ?", roleId).Delete(&rm).Error; err != nil { + return err + } + var role SysRole + if err := tx.Table("sys_role").Where("role_id = ?", roleId).First(&role).Error; err != nil { + return err + } + sql3 := "delete from sys_casbin_rule where v0= '" + role.RoleKey + "';" + if err := tx.Exec(sql3).Error; err != nil { + return err + } + return nil + +} + +// 该方法即将弃用 +func (rm *RoleMenu) BatchDeleteRoleMenu(tx *gorm.DB, roleIds []int) error { + if err := tx.Table("sys_role_menu").Where("role_id in (?)", roleIds).Delete(&rm).Error; err != nil { + return err + } + var role []SysRole + if err := tx.Table("sys_role").Where("role_id in (?)", roleIds).Find(&role).Error; err != nil { + return err + } + sql := "" + for i := 0; i < len(role); i++ { + sql += "delete from sys_casbin_rule where v0= '" + role[i].RoleName + "';" + } + if err := tx.Exec(sql).Error; err != nil { + return err + } + if err := tx.Commit().Error; err != nil { + return err + } + return nil + +} + +func (rm *RoleMenu) Insert(tx *gorm.DB, enforcer *casbin.SyncedEnforcer, roleId int, menuId []int) error { + var err error + var ( + role SysRole + menu []models.Menu + casbinRules []models.CasbinRule // casbinRule 待插入队列 + ) + // 在事务中做一些数据库操作(从这一点使用'tx',而不是'db') + if err = tx.Table("sys_role").Where("role_id = ?", roleId).First(&role).Error; err != nil { + return err + } + if err = tx.Table("sys_menu").Where("menu_id in (?)", menuId).Find(&menu).Error; err != nil { + return err + } + //ORM不支持批量插入所以需要拼接 sql 串 + sysRoleMenuSql := "INSERT INTO `sys_role_menu` (`role_id`,`menu_id`,`role_name`) VALUES " + + for i, m := range menu { + // 拼装'role_menu'表批量插入SQL语句 + sysRoleMenuSql += fmt.Sprintf("(%d,%d,'%s')", role.RoleId, m.MenuId, role.RoleKey) + if i == len(menu)-1 { + sysRoleMenuSql += ";" //最后一条数据 以分号结尾 + } else { + sysRoleMenuSql += "," + } + if m.MenuType == "A" { + // 加入队列 + casbinRules = append(casbinRules, + models.CasbinRule{ + V0: role.RoleKey, + V1: m.Path, + V2: m.Action, + }) + } + } + // 执行批量插入sys_role_menu + if err = tx.Exec(sysRoleMenuSql).Error; err != nil { + return err + } + + // 执行批量插入sys_casbin_rule + if len(casbinRules) > 0 { + if err = tx.Create(&casbinRules).Error; err != nil { + return err + } + } + return nil +} + +func (rm *RoleMenu) Delete(tx *gorm.DB, RoleId string, MenuID string) (bool, error) { + rm.RoleId, _ = tools.StringToInt(RoleId) + table := tx.Table("sys_role_menu").Where("role_id = ?", RoleId) + if MenuID != "" { + table = table.Where("menu_id = ?", MenuID) + } + if err := table.Delete(&rm).Error; err != nil { + return false, err + } + return true, nil + +} diff --git a/app/admin/router/initrouter.go b/app/admin/router/initrouter.go index 8bc82ed7..c3b62452 100644 --- a/app/admin/router/initrouter.go +++ b/app/admin/router/initrouter.go @@ -32,7 +32,6 @@ func InitRouter() { if config.SslConfig.Enable { r.Use(handler.TlsHandler()) } - r.Use(middleware.WithContextDb(middleware.GetGormFromConfig(global.Cfg))) r.Use(common.Sentinel()) middleware.InitMiddleware(r) diff --git a/app/admin/router/sys_role.go b/app/admin/router/sys_role.go index 189c0e6c..809580f1 100644 --- a/app/admin/router/sys_role.go +++ b/app/admin/router/sys_role.go @@ -22,4 +22,5 @@ func registerSysRoleRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddle r.PUT("/:id", api.UpdateSysRole) r.DELETE("/:id", api.DeleteSysRole) } + v1.PUT("/roledatascope", api.UpdateRoleDataScope) } diff --git a/app/admin/router/sysrouter.go b/app/admin/router/sysrouter.go index aaf96b0a..b3d55e04 100644 --- a/app/admin/router/sysrouter.go +++ b/app/admin/router/sysrouter.go @@ -117,7 +117,7 @@ func registerBaseRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddlewar { v1auth.GET("/getinfo", system.GetInfo) - v1auth.PUT("/roledatascope", system.UpdateRoleDataScope) + //v1auth.PUT("/roledatascope", system.UpdateRoleDataScope) v1auth.GET("/roleMenuTreeselect/:roleId", system.GetMenuTreeRoleselect) v1auth.GET("/roleDeptTreeselect/:roleId", system.GetDeptTreeRoleselect) @@ -168,15 +168,15 @@ func registerPostRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddlewar // } //} -func registerRoleRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddleware) { - role := v1.Group("/role").Use(authMiddleware.MiddlewareFunc()).Use(middleware.AuthCheckRole()) - { - role.GET("/:roleId", system.GetRole) - role.POST("", system.InsertRole) - role.PUT("", system.UpdateRole) - role.DELETE("/:roleId", system.DeleteRole) - } -} +//func registerRoleRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddleware) { +// role := v1.Group("/role").Use(authMiddleware.MiddlewareFunc()).Use(middleware.AuthCheckRole()) +// { +// role.GET("/:roleId", system.GetRole) +// role.POST("", system.InsertRole) +// role.PUT("", system.UpdateRole) +// role.DELETE("/:roleId", system.DeleteRole) +// } +//} func registerSysUserRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddleware) { sysuser := v1.Group("/sysUser").Use(authMiddleware.MiddlewareFunc()).Use(middleware.AuthCheckRole()) diff --git a/app/admin/service/dto/sys_role.go b/app/admin/service/dto/sys_role.go index 40e5022f..f5c01317 100644 --- a/app/admin/service/dto/sys_role.go +++ b/app/admin/service/dto/sys_role.go @@ -51,6 +51,7 @@ type SysRoleControl struct { Remark string `form:"remark" comment:"备注"` // 备注 Admin bool `form:"admin" comment:"是否管理员"` DataScope string `form:"dataScope" comment:"是否管理员"` + MenuIds []int `json:"menuIds"` } // Bind 映射上下文中的结构体数据 @@ -86,6 +87,7 @@ func (s *SysRoleControl) Generate() (*system.SysRole, error) { Remark: s.Remark, Admin: s.Admin, DataScope: s.DataScope, + MenuIds: s.MenuIds, }, nil } @@ -131,4 +133,5 @@ func (s *SysRoleById) GenerateM() (*models.SysRole, error) { type RoleDataScopeReq struct { RoleId int `json:"roleId" binding:"required"` DataScope string `json:"dataScope" binding:"required"` + DeptIds []int `json:"deptIds"` } diff --git a/app/admin/service/sys_role.go b/app/admin/service/sys_role.go index 599661d8..00692ff4 100644 --- a/app/admin/service/sys_role.go +++ b/app/admin/service/sys_role.go @@ -2,12 +2,16 @@ package service import ( "errors" + + "gorm.io/gorm" + + "go-admin/app/admin/models" "go-admin/app/admin/models/system" "go-admin/app/admin/service/dto" cDto "go-admin/common/dto" + orm "go-admin/common/global" "go-admin/common/log" "go-admin/common/service" - "gorm.io/gorm" ) type SysRole struct { @@ -48,35 +52,66 @@ func (e *SysRole) GetSysRole(d *dto.SysRoleById, model *system.SysRole) error { log.Errorf("msgID[%s] db error:%s", msgID, err) return err } - if db.Error != nil { + if err != nil { log.Errorf("msgID[%s] db error:%s", msgID, err) return err } + data.MenuIds, err = e.GetRoleMenuId(data.RoleId) + if err != nil { + log.Errorf("msgID[%s] get menuIds error, %s", msgID, err.Error()) + return err + } return nil } // InsertSysRole 创建SysRole对象 -func (e *SysRole) InsertSysRole(model *system.SysRole) error { +func (e *SysRole) InsertSysRole(c *system.SysRole) error { var err error var data system.SysRole msgID := e.MsgID - err = e.Orm.Model(&data). - Create(model).Error + tx := e.Orm.Begin() + defer func() { + if err != nil { + tx.Rollback() + } else { + tx.Commit() + } + }() + + err = tx.Model(&data). + Create(c).Error if err != nil { log.Errorf("msgID[%s] db error:%s", msgID, err) return err } + if len(c.MenuIds) > 0 { + s := SysRoleMenu{} + s.Orm = e.Orm + s.MsgID = msgID + err = s.ReloadRule(tx, c.RoleId, c.MenuIds) + if err != nil { + log.Errorf("msgID[%s] reload casbin rule error, %", msgID, err.Error()) + return err + } + } return nil } // UpdateSysRole 修改SysRole对象 func (e *SysRole) UpdateSysRole(c *system.SysRole) error { var err error - var data system.SysRole msgID := e.MsgID - db := e.Orm.Model(&data). + tx := e.Orm.Debug().Begin() + defer func() { + if err != nil { + tx.Rollback() + } else { + tx.Commit() + } + }() + db := tx.Model(&c). Where(c.GetId()).Updates(c) if db.Error != nil { log.Errorf("msgID[%s] db error:%s", msgID, err) @@ -85,6 +120,23 @@ func (e *SysRole) UpdateSysRole(c *system.SysRole) error { if db.RowsAffected == 0 { return errors.New("无权更新该数据") } + + var t system.RoleMenu + err = t.DeleteRoleMenu(tx, c.RoleId) + if err != nil { + log.Errorf("msgID[%s] delete role menu error, %", msgID, err.Error()) + return err + } + if len(c.MenuIds) > 0 { + s := SysRoleMenu{} + s.Orm = e.Orm + s.MsgID = msgID + err = s.ReloadRule(tx, c.RoleId, c.MenuIds) + if err != nil { + log.Errorf("msgID[%s] reload casbin rule error, %", msgID, err.Error()) + return err + } + } return nil } @@ -94,7 +146,16 @@ func (e *SysRole) RemoveSysRole(d *dto.SysRoleById) error { var data system.SysRole msgID := e.MsgID - db := e.Orm.Model(&data).Delete(&data, d.GetId()) + tx := e.Orm.Begin() + defer func() { + if err != nil { + tx.Rollback() + } else { + tx.Commit() + } + }() + + db := tx.Model(&data).Delete(&data, d.GetId()) if db.Error != nil { err = db.Error log.Errorf("MsgID[%s] Delete error: %s", msgID, err) @@ -104,5 +165,67 @@ func (e *SysRole) RemoveSysRole(d *dto.SysRoleById) error { err = errors.New("无权删除该数据") return err } + s := SysRoleMenu{} + s.Orm = db + s.MsgID = msgID + err = s.DeleteRoleMenu(tx, d.Id) + if err != nil { + log.Errorf("msgID[%s] insert role menu error, %", msgID, err.Error()) + return err + } return nil } + +// 获取角色对应的菜单ids +func (e *SysRole) GetRoleMenuId(roleId int) ([]int, error) { + menuIds := make([]int, 0) + menuList := make([]models.MenuIdList, 0) + if err := orm.Eloquent.Table("sys_role_menu"). + Select("sys_role_menu.menu_id"). + Where("role_id = ? ", roleId). + Where(" sys_role_menu.menu_id not in(select sys_menu.parent_id from sys_role_menu "+ + "LEFT JOIN sys_menu on sys_menu.menu_id=sys_role_menu.menu_id where role_id =? and parent_id is not null)", roleId). + Find(&menuList).Error; err != nil { + return nil, err + } + + for i := 0; i < len(menuList); i++ { + menuIds = append(menuIds, menuList[i].MenuId) + } + return menuIds, nil +} + +func (e *SysRole) UpdateDataScope(c *system.SysRole) (err error) { + tx := e.Orm.Begin() + defer func() { + if err != nil { + tx.Rollback() + } else { + tx.Commit() + } + }() + err = tx.Model(&system.SysRole{}).Where("role_id = ?", c.RoleId).Select("data_scope, update_by").Updates(c).Error + if err != nil { + return err + } + + err = tx.Where("role_id = ?", c.RoleId).Delete(&system.SysRoleDept{}).Error + if err != nil { + return err + } + + if c.DataScope == "2" { + deptRoles := make([]system.SysRoleDept, len(c.DeptIds)) + for i := range c.DeptIds { + deptRoles[i] = system.SysRoleDept{ + RoleId: c.RoleId, + DeptId: c.DeptIds[i], + } + } + err = tx.Create(&deptRoles).Error + if err != nil { + return err + } + } + return err +} diff --git a/app/admin/service/sys_role_menu.go b/app/admin/service/sys_role_menu.go new file mode 100644 index 00000000..a4fe849e --- /dev/null +++ b/app/admin/service/sys_role_menu.go @@ -0,0 +1,113 @@ +package service + +import ( + "go-admin/app/admin/models" + "go-admin/app/admin/models/system" + "go-admin/common/log" + "go-admin/common/service" + "gorm.io/gorm" +) + +type SysRoleMenu struct { + service.Service +} + +func (e *SysRoleMenu) ReloadRule(tx *gorm.DB, roleId int, menuId []int) (err error) { + var role system.SysRole + + msgID := e.MsgID + + menu := make([]models.Menu, 0) + roleMenu := make([]system.RoleMenu, len(menuId)) + casbinRule := make([]models.CasbinRule, 0) + //先删除所有的 + err = e.DeleteRoleMenu(tx, roleId) + if err != nil { + return + } + + // 在事务中做一些数据库操作(从这一点使用'tx',而不是'db') + err = tx.Where("role_id = ?", roleId).First(&role).Error + if err != nil { + log.Errorf("msgID[%s] get role error, %s", msgID, err.Error()) + return + } + err = tx.Where("menu_id in (?)", menuId). + //Select("path, action, menu_id, menu_type"). + Find(&menu).Error + if err != nil { + log.Errorf("msgID[%s] get menu error, %s", msgID, err.Error()) + return + } + for i := range menu { + roleMenu[i] = system.RoleMenu{ + RoleId: role.RoleId, + MenuId: menu[i].MenuId, + RoleName: role.RoleKey, + } + if menu[i].MenuType == "A" { + casbinRule = append(casbinRule, models.CasbinRule{ + PType: "p", + V0: role.RoleKey, + V1: menu[i].Path, + V2: menu[i].Action, + }) + } + } + err = tx.Create(&roleMenu).Error + if err != nil { + log.Errorf("msgID[%s] batch create role's menu error, %s", msgID, err.Error()) + return + } + if len(casbinRule) > 0 { + err = tx.Create(&casbinRule).Error + if err != nil { + log.Errorf("msgID[%s] batch create casbin rule error, %s", msgID, err.Error()) + return + } + } + + return +} + +func (e *SysRoleMenu) DeleteRoleMenu(tx *gorm.DB, roleId int) (err error) { + msgID := e.MsgID + err = tx.Where("role_id = ?", roleId). + Delete(&system.SysRoleDept{}).Error + if err != nil { + log.Errorf("msgID[%s] delete role's dept error, %s", msgID, err.Error()) + return + } + err = tx.Where("role_id = ?", roleId). + Delete(&system.RoleMenu{}).Error + if err != nil { + log.Errorf("msgID[%s] delete role's menu error, %s", msgID, err.Error()) + return + } + var role system.SysRole + err = tx.Where("role_id = ?", roleId). + First(&role).Error + if err != nil { + log.Errorf("msgID[%s] get role error, %s", msgID, err.Error()) + return + } + err = tx.Where("v0 = ?", role.RoleKey). + Delete(&models.CasbinRule{}).Error + if err != nil { + log.Errorf("msgID[%s] delete casbin rule error, %s", msgID, err.Error()) + return + } + return +} + +func (e *SysRoleMenu) GetIDS(tx *gorm.DB, roleName string) ([]system.MenuPath, error) { + var r []system.MenuPath + table := tx.Select("sys_menu.path").Table("sys_role_menu") + table = table.Joins("left join sys_role on sys_role.role_id=sys_role_menu.role_id") + table = table.Joins("left join sys_menu on sys_menu.id=sys_role_menu.menu_id") + table = table.Where("sys_role.role_name = ? and sys_menu.type=1", roleName) + if err := table.Find(&r).Error; err != nil { + return nil, err + } + return r, nil +} diff --git a/cmd/api/server.go b/cmd/api/server.go index 55bdffda..3197a064 100644 --- a/cmd/api/server.go +++ b/cmd/api/server.go @@ -17,19 +17,14 @@ import ( "go-admin/common/database" "go-admin/common/global" "go-admin/common/log" - mycasbin "go-admin/pkg/casbin" "go-admin/pkg/logger" "go-admin/tools" "go-admin/tools/config" - "go-admin/tools/trace" ) var ( - configYml string - port string - mode string - traceStart bool - StartCmd = &cobra.Command{ + configYml string + StartCmd = &cobra.Command{ Use: "server", Short: "Start API server", Example: "go-admin server -c config/settings.yml", @@ -47,9 +42,6 @@ var AppRouters = make([]func(), 0) func init() { StartCmd.PersistentFlags().StringVarP(&configYml, "config", "c", "config/settings.yml", "Start server with provided configuration file") - StartCmd.PersistentFlags().StringVarP(&port, "port", "p", "8000", "Tcp port server listening on") - StartCmd.PersistentFlags().StringVarP(&mode, "mode", "m", "dev", "server mode ; eg:dev,test,prod") - StartCmd.PersistentFlags().BoolVarP(&traceStart, "traceStart", "t", false, "start traceStart app dash") //注册路由 fixme 其他应用的路由,在本目录新建文件放在init方法 AppRouters = append(AppRouters, router.InitRouter) @@ -65,9 +57,7 @@ func setup() { global.JobLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "job") global.RequestLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "request") //3. 初始化数据库链接 - database.Setup(config.DatabaseConfig.Driver) - //4. 接口访问控制加载 - global.CasbinEnforcer = mycasbin.Setup(global.Eloquent, "sys_") + database.Setup() usageStr := `starting api server` log.Info(usageStr) @@ -85,7 +75,7 @@ func run() error { engine = gin.New() } - if mode == "dev" { + if config.ApplicationConfig.Mode == "dev" { //监控 AppRouters = append(AppRouters, router.Monitor) } @@ -107,12 +97,6 @@ func run() error { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - if traceStart { - //链路追踪, fixme 页面显示需要自备梯子 - trace.Start() - defer trace.Stop(ctx) - } - go func() { // 服务连接 if config.SslConfig.Enable { diff --git a/cmd/config/server.go b/cmd/config/server.go index f23cba39..38d65afa 100644 --- a/cmd/config/server.go +++ b/cmd/config/server.go @@ -12,7 +12,6 @@ import ( var ( configYml string - mode string StartCmd = &cobra.Command{ Use: "config", Short: "Get Application config info", @@ -42,7 +41,8 @@ func run() { } fmt.Println("jwt:", string(jwt)) - database, errs := json.MarshalIndent(config.DatabaseConfig, "", " ") //转换成JSON返回的是byte[] + // todo 需要兼容 + database, errs := json.MarshalIndent(config.DatabasesConfig, "", " ") //转换成JSON返回的是byte[] if errs != nil { fmt.Println(errs.Error()) } diff --git a/cmd/migrate/migration/version/1599190683659_tables.go b/cmd/migrate/migration/version/1599190683659_tables.go index 408ee5ce..ba500d39 100644 --- a/cmd/migrate/migration/version/1599190683659_tables.go +++ b/cmd/migrate/migration/version/1599190683659_tables.go @@ -28,7 +28,7 @@ func _1599190683659Tables(db *gorm.DB, version string) error { new(system.SysLoginLog), new(system.SysOperaLog), new(models.RoleMenu), - new(models.SysRoleDept), + new(system.SysRoleDept), new(models.SysUser), new(system.SysRole), new(models.Post), diff --git a/cmd/migrate/server.go b/cmd/migrate/server.go index 99da7ead..f47ebf73 100644 --- a/cmd/migrate/server.go +++ b/cmd/migrate/server.go @@ -72,7 +72,7 @@ func migrateModel() error { } func initDB() error { //3. 初始化数据库链接 - database.Setup(config.DatabaseConfig.Driver) + database.Setup() //4. 数据库迁移 fmt.Println("数据库迁移开始") _ = migrateModel() diff --git a/common/config/config.go b/common/config/config.go index 5222d114..0b5fa438 100644 --- a/common/config/config.go +++ b/common/config/config.go @@ -1,47 +1,47 @@ package config import ( - "database/sql" "net/http" + "github.com/casbin/casbin/v2" "github.com/go-admin-team/go-admin-core/logger" + "gorm.io/gorm" ) type Config struct { - saas bool - dbs map[string]*DBConfig - db *DBConfig - engine http.Handler + dbs map[string]*gorm.DB + casbins map[string]*casbin.SyncedEnforcer + engine http.Handler } -type DBConfig struct { - Driver string - DB *sql.DB -} - -// SetDbs 设置对应key的db -func (c *Config) SetDbs(key string, db *DBConfig) { +// SetDb 设置对应key的db +func (c *Config) SetDb(key string, db *gorm.DB) { c.dbs[key] = db } -// GetDbs 获取所有map里的db数据 -func (c *Config) GetDbs() map[string]*DBConfig { +// GetDb 获取所有map里的db数据 +func (c *Config) GetDb() map[string]*gorm.DB { return c.dbs } // GetDbByKey 根据key获取db -func (c *Config) GetDbByKey(key string) *DBConfig { +func (c *Config) GetDbByKey(key string) *gorm.DB { + if db, ok := c.dbs["*"]; ok { + return db + } return c.dbs[key] } -// SetDb 设置单个db -func (c *Config) SetDb(db *DBConfig) { - c.db = db +func (c *Config) SetCasbin(key string, enforcer *casbin.SyncedEnforcer) { + c.casbins[key] = enforcer } -// GetDb 获取单个db -func (c *Config) GetDb() *DBConfig { - return c.db +// GetCasbinKey 根据key获取casbin +func (c *Config) GetCasbinKey(key string) *casbin.SyncedEnforcer { + if e, ok := c.casbins["*"]; ok { + return e + } + return c.casbins[key] } // SetEngine 设置路由引擎 @@ -64,16 +64,10 @@ 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{} +// NewConfig 默认值 +func NewConfig() *Config { + return &Config{ + dbs: make(map[string]*gorm.DB), + casbins: make(map[string]*casbin.SyncedEnforcer), + } } diff --git a/common/config/type.go b/common/config/type.go index a8060a62..c3f772a1 100644 --- a/common/config/type.go +++ b/common/config/type.go @@ -1,22 +1,21 @@ package config import ( + "github.com/casbin/casbin/v2" "net/http" "github.com/go-admin-team/go-admin-core/logger" + "gorm.io/gorm" ) 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(key string, db *gorm.DB) + GetDb() map[string]*gorm.DB + GetDbByKey(key string) *gorm.DB - //单库业务实现这两个接口 - SetDb(db *DBConfig) - GetDb() *DBConfig + SetCasbin(key string, enforcer *casbin.SyncedEnforcer) + GetCasbinKey(key string) *casbin.SyncedEnforcer //使用的路由 SetEngine(engine http.Handler) diff --git a/common/database/initialize.go b/common/database/initialize.go index f59ac3bb..0a188546 100644 --- a/common/database/initialize.go +++ b/common/database/initialize.go @@ -1,17 +1,59 @@ -// +build !sqlite3 - package database -// Setup 配置数据库 -func Setup(driver string) { - dbType := driver - if dbType == "mysql" { - var db = new(Mysql) - db.Setup() - } +import ( + . "log" + "time" - if dbType == "postgres" { - var db = new(PgSql) - db.Setup() + logCore "github.com/go-admin-team/go-admin-core/logger" + "gorm.io/gorm" + "gorm.io/gorm/logger" + "gorm.io/gorm/schema" + + "go-admin/common/global" + "go-admin/common/log" + mycasbin "go-admin/pkg/casbin" + "go-admin/tools" + toolsConfig "go-admin/tools/config" +) + +// Setup 配置数据库 +func Setup() { + for k := range toolsConfig.DatabasesConfig { + setupSimpleDatabase(k, toolsConfig.DatabasesConfig[k]) } } + +func setupSimpleDatabase(host string, c *toolsConfig.Database) { + if global.Driver == "" { + global.Driver = c.Driver + } + log.Infof("%s => %s", host, tools.Green(c.Source)) + db, err := gorm.Open(open[c.Driver](c.Source), &gorm.Config{ + NamingStrategy: schema.NamingStrategy{ + SingularTable: true, + }, + Logger: logger.New( + New(logCore.DefaultLogger.Options().Out, "\r\n", LstdFlags), + logger.Config{ + SlowThreshold: time.Second, + Colorful: true, + LogLevel: logger.LogLevel( + logCore.DefaultLogger.Options().Level.LevelForGorm()), + }, + ), + }) + if err != nil { + log.Fatal(tools.Red(c.Driver+" connect error :"), err) + } else { + log.Info(tools.Green(c.Driver + " connect success !")) + } + + e := mycasbin.Setup(db, "sys_") + + if host == "*" { + global.Eloquent = db + } + + global.Cfg.SetDb(host, db) + global.Cfg.SetCasbin(host, e) +} diff --git a/common/database/initialize_sqlite3.go b/common/database/initialize_sqlite3.go deleted file mode 100644 index 7bd6523e..00000000 --- a/common/database/initialize_sqlite3.go +++ /dev/null @@ -1,21 +0,0 @@ -// +build sqlite3 - -package database - -func Setup(driver string) { - dbType := driver - if dbType == "mysql" { - var db = new(Mysql) - db.Setup() - } - - if dbType == "sqlite3" { - var db = new(SqLite) - db.Setup() - } - - if dbType == "postgres" { - var db = new(PgSql) - db.Setup() - } -} diff --git a/common/database/interface.go b/common/database/interface.go deleted file mode 100644 index 59c5e1ce..00000000 --- a/common/database/interface.go +++ /dev/null @@ -1,11 +0,0 @@ -package database - -import "gorm.io/gorm" - -// Database 数据库配置 -type Database interface { - Setup() - Open(conn string, cfg *gorm.Config) (db *gorm.DB, err error) - GetConnect() string - GetDriver() string -} diff --git a/common/database/mysql_drive.go b/common/database/mysql_drive.go deleted file mode 100644 index faa4d977..00000000 --- a/common/database/mysql_drive.go +++ /dev/null @@ -1,77 +0,0 @@ -package database - -import ( - "database/sql" - . "log" - "time" - - goAdminLogger "github.com/go-admin-team/go-admin-core/logger" - "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/common/log" - "go-admin/tools" - toolsConfig "go-admin/tools/config" -) - -// Mysql mysql配置结构体 -type Mysql struct { -} - -// Setup 配置步骤 -func (e *Mysql) Setup() { - global.Source = e.GetConnect() - log.Info(tools.Green(global.Source)) - db, err := sql.Open("mysql", global.Source) - if err != nil { - log.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, - }, - }) - if err != nil { - log.Fatal(tools.Red(e.GetDriver()+" connect error :"), err) - } else { - log.Info(tools.Green(e.GetDriver() + " connect success !")) - } - - if global.Eloquent.Error != nil { - log.Fatal(tools.Red(" database error :"), global.Eloquent.Error) - } - - if toolsConfig.LoggerConfig.EnabledDB { - global.Eloquent.Logger = logger.New( - New(goAdminLogger.DefaultLogger.Options().Out, "\r\n", LstdFlags), - logger.Config{ - SlowThreshold: time.Second, - Colorful: true, - LogLevel: logger.LogLevel( - goAdminLogger.DefaultLogger.Options().Level.LevelForGorm()), - }) - } -} - -// Open 打开数据库连接 -func (e *Mysql) Open(db *sql.DB, cfg *gorm.Config) (*gorm.DB, error) { - return gorm.Open(mysql.New(mysql.Config{Conn: db}), cfg) -} - -// GetConnect 获取数据库连接 -func (e *Mysql) GetConnect() string { - return toolsConfig.DatabaseConfig.Source -} - -// GetDriver 获取连接 -func (e *Mysql) GetDriver() string { - return toolsConfig.DatabaseConfig.Driver -} diff --git a/common/database/open.go b/common/database/open.go new file mode 100644 index 00000000..1dbd4a32 --- /dev/null +++ b/common/database/open.go @@ -0,0 +1,14 @@ +// +build !sqlite3 + +package database + +import ( + "gorm.io/driver/mysql" + "gorm.io/driver/postgres" + "gorm.io/gorm" +) + +var open = map[string]func(string) gorm.Dialector{ + "mysql": mysql.Open, + "postgres": postgres.Open, +} diff --git a/common/database/open_sqlite3.go b/common/database/open_sqlite3.go new file mode 100644 index 00000000..387d1286 --- /dev/null +++ b/common/database/open_sqlite3.go @@ -0,0 +1,16 @@ +// +build sqlite3 + +package database + +import ( + "gorm.io/driver/mysql" + "gorm.io/driver/postgres" + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +var open = map[string]func(string) gorm.Dialector{ + "mysql": mysql.Open, + "postgres": postgres.Open, + "sqlite3": sqlite.Open, +} diff --git a/common/database/pgsql_driver.go b/common/database/pgsql_driver.go deleted file mode 100644 index b1d9bead..00000000 --- a/common/database/pgsql_driver.go +++ /dev/null @@ -1,75 +0,0 @@ -package database - -import ( - "database/sql" - "go-admin/common/log" - . "log" - "time" - - "gorm.io/driver/postgres" - "gorm.io/gorm" - "gorm.io/gorm/logger" - "gorm.io/gorm/schema" - - goAdminLogger "github.com/go-admin-team/go-admin-core/logger" - "go-admin/common/config" - "go-admin/common/global" - "go-admin/tools" - toolsConfig "go-admin/tools/config" -) - -type PgSql struct { -} - -func (e *PgSql) Setup() { - var err error - - global.Source = e.GetConnect() - log.Info(global.Source) - 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, - }, - }) - if err != nil { - log.Fatalf("%s connect error %v", e.GetDriver(), err) - } else { - log.Infof("%s connect success!", e.GetDriver()) - } - - if global.Eloquent.Error != nil { - log.Fatalf("database error %v", global.Eloquent.Error) - } - - if toolsConfig.LoggerConfig.EnabledDB { - global.Eloquent.Logger = logger.New( - New(goAdminLogger.DefaultLogger.Options().Out, "\r\n", LstdFlags), - logger.Config{ - SlowThreshold: time.Second, - Colorful: true, - LogLevel: logger.LogLevel( - goAdminLogger.DefaultLogger.Options().Level.LevelForGorm()), - }) - } -} - -// 打开数据库连接 -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 toolsConfig.DatabaseConfig.Source -} - -func (e *PgSql) GetDriver() string { - return toolsConfig.DatabaseConfig.Driver -} diff --git a/common/database/sqlite3_driver.go b/common/database/sqlite3_driver.go deleted file mode 100644 index c7fb5173..00000000 --- a/common/database/sqlite3_driver.go +++ /dev/null @@ -1,77 +0,0 @@ -// +build sqlite3 - -package database - -import ( - "database/sql" - . "log" - "os" - "time" - - "gorm.io/driver/sqlite" - "gorm.io/gorm" - "gorm.io/gorm/logger" - "gorm.io/gorm/schema" - - "go-admin/common/config" - "go-admin/common/global" - "go-admin/common/log" - "go-admin/tools" - toolsConfig "go-admin/tools/config" -) - -type SqLite struct { -} - -func (e *SqLite) Setup() { - var err error - - global.Source = e.GetConnect() - log.Info(global.Source) - db, err := sql.Open("sqlite3", global.Source) - if err != nil { - global.Logger.Fatal(tools.Red(e.GetDriver()+" connect error :"), err) - } - global.Cfg.SetDb(&config.DBConfig{ - Driver: "sqlite3", - DB: db, - }) - global.Eloquent, err = e.Open(e.GetConnect(), &gorm.Config{ - NamingStrategy: schema.NamingStrategy{ - SingularTable: true, - }, - }) - - if err != nil { - log.Fatalf("%s connect error %v", e.GetDriver(), err) - } else { - log.Infof("%s connect success!", e.GetDriver()) - } - - if global.Eloquent.Error != nil { - log.Fatalf("database error %v", global.Eloquent.Error) - } - - if toolsConfig.LoggerConfig.EnabledDB { - global.Eloquent.Logger = logger.New( - New(os.Stdout, "\r\n", LstdFlags), logger.Config{ - SlowThreshold: time.Second, - Colorful: true, - LogLevel: logger.Info, - }, - ) - } -} - -// 打开数据库连接 -func (*SqLite) Open(conn string, cfg *gorm.Config) (db *gorm.DB, err error) { - return gorm.Open(sqlite.Open(conn), cfg) -} - -func (e *SqLite) GetConnect() string { - return toolsConfig.DatabaseConfig.Source -} - -func (e *SqLite) GetDriver() string { - return toolsConfig.DatabaseConfig.Driver -} diff --git a/common/dto/search.go b/common/dto/search.go index 5c2cbcf8..72cc60e4 100644 --- a/common/dto/search.go +++ b/common/dto/search.go @@ -2,9 +2,8 @@ package dto import ( "github.com/go-admin-team/go-admin-core/tools/search" + "go-admin/common/global" "gorm.io/gorm" - - "go-admin/tools/config" ) type GeneralDelDto struct { @@ -45,7 +44,7 @@ func MakeCondition(q interface{}) func(db *gorm.DB) *gorm.DB { GormPublic: search.GormPublic{}, Join: make([]*search.GormJoin, 0), } - search.ResolveSearchQuery(config.DatabaseConfig.Driver, q, condition) + search.ResolveSearchQuery(global.Driver, q, condition) for _, join := range condition.Join { if join == nil { continue diff --git a/common/global/adm.go b/common/global/adm.go index 291e4fef..58d20b09 100644 --- a/common/global/adm.go +++ b/common/global/adm.go @@ -1,7 +1,6 @@ package global import ( - "github.com/casbin/casbin/v2" "github.com/gin-gonic/gin" "github.com/robfig/cron/v3" "gorm.io/gorm" @@ -15,10 +14,9 @@ const ( Version = "1.2.3" ) -var Cfg config.Conf = config.DefaultConfig() +var Cfg config.Conf = config.NewConfig() var GinEngine *gin.Engine -var CasbinEnforcer *casbin.SyncedEnforcer var Eloquent *gorm.DB var GADMCron *cron.Cron diff --git a/common/global/casbin.go b/common/global/casbin.go index bbe300fe..f10d9df5 100644 --- a/common/global/casbin.go +++ b/common/global/casbin.go @@ -3,13 +3,17 @@ package global import ( "github.com/casbin/casbin/v2" "github.com/casbin/casbin/v2/log" + "github.com/gin-gonic/gin" + + "go-admin/tools" ) -func LoadPolicy() (*casbin.SyncedEnforcer, error) { - if err := CasbinEnforcer.LoadPolicy(); err == nil { - return CasbinEnforcer, err +func LoadPolicy(c *gin.Context) (*casbin.SyncedEnforcer, error) { + if err := Cfg.GetCasbinKey(c.Request.Host).LoadPolicy(); err == nil { + return Cfg.GetCasbinKey(c.Request.Host), err } else { - log.LogPrintf("casbin rbac_model or policy init error, message: %v \r\n", err.Error()) + msgID := tools.GenerateMsgIDFromContext(c) + log.LogPrintf("msgID[%s] casbin rbac_model or policy init error, message: %v \r\n", msgID, err.Error()) return nil, err } } diff --git a/common/middleware/db.go b/common/middleware/db.go index 272ec3df..f2d3f895 100644 --- a/common/middleware/db.go +++ b/common/middleware/db.go @@ -2,16 +2,11 @@ package middleware import ( "github.com/gin-gonic/gin" - "gorm.io/gorm" + + "go-admin/common/global" ) -func WithContextDb(dbMap map[string]*gorm.DB) gin.HandlerFunc { - return func(c *gin.Context) { - if db, ok := dbMap["*"]; ok { - c.Set("db", db) - } else { - c.Set("db", dbMap[c.Request.Host]) - } - c.Next() - } +func WithContextDb(c *gin.Context) { + c.Set("db", global.Cfg.GetDbByKey(c.Request.Host)) + c.Next() } diff --git a/examples/run.go b/examples/run.go index a6bbc3dc..74a23dec 100644 --- a/examples/run.go +++ b/examples/run.go @@ -16,12 +16,11 @@ import ( ) func main() { - var err error - global.Eloquent, err = gorm.Open(mysql.Open("root:123456@tcp/inmg?charset=utf8&parseTime=True&loc=Local"), &gorm.Config{}) + db, err := gorm.Open(mysql.Open("root:123456@tcp/inmg?charset=utf8&parseTime=True&loc=Local"), &gorm.Config{}) if err != nil { panic(err) } - global.CasbinEnforcer = mycasbin.Setup(global.Eloquent, "sys_") + _ = mycasbin.Setup(db, "sys_") global.Logger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "bus") global.JobLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "job") global.RequestLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "request") diff --git a/tools/config/config.go b/tools/config/config.go index 6edee1f7..2a97eb95 100644 --- a/tools/config/config.go +++ b/tools/config/config.go @@ -9,6 +9,8 @@ import ( var ( ExtendConfig interface{} + _watch config.Watcher + _cfg *Settings ) // Settings 兼容原先的配置结构 @@ -18,19 +20,25 @@ type Settings struct { // Config 配置集合 type Config struct { - Application *Application `yaml:"application"` - Ssl *Ssl `yaml:"ssl"` - Logger *Logger `yaml:"logger"` - Jwt *Jwt `yaml:"jwt"` - Database *Database `yaml:"database"` - Gen *Gen `yaml:"gen"` - Extend interface{} `yaml:"extend"` + Application *Application `yaml:"application"` + Ssl *Ssl `yaml:"ssl"` + Logger *Logger `yaml:"logger"` + Jwt *Jwt `yaml:"jwt"` + Database *Database `yaml:"database"` + Databases *map[string]*Database `yaml:"databases"` + Gen *Gen `yaml:"gen"` + Extend interface{} `yaml:"extend"` } -var ( - _watch config.Watcher - _cfg *Settings -) +// 多db改造 +func (e *Config) multiDatabase() { + if len(*e.Databases) == 0 { + *e.Databases = map[string]*Database{ + "*": e.Database, + } + + } +} // Setup 载入配置文件 func Setup(f func(opts ...source.Option) source.Source, options ...source.Option) { @@ -50,7 +58,8 @@ func Setup(f func(opts ...source.Option) source.Source, options ...source.Option Ssl: SslConfig, Logger: LoggerConfig, Jwt: JwtConfig, - Database: DatabaseConfig, + Database: new(Database), + Databases: &DatabasesConfig, Gen: GenConfig, Extend: ExtendConfig, }} @@ -58,6 +67,7 @@ func Setup(f func(opts ...source.Option) source.Source, options ...source.Option if err != nil { log.Fatal(fmt.Sprintf("Scan config fail: %s", err.Error())) } + _cfg.Settings.multiDatabase() _watch, err = c.Watch() if err != nil { @@ -65,7 +75,7 @@ func Setup(f func(opts ...source.Option) source.Source, options ...source.Option } } -// Watch 配置监听, 重载时报错,不影响运行 +// Watch 配置监听, 重载时报错,不影响运行 fixme 数据连接 redis连接还没支持动态配置 func Watch() { for { v, err := _watch.Next() @@ -80,7 +90,7 @@ func Watch() { log.Println(fmt.Sprintf("Scan config fail: %s", err.Error())) break } - fmt.Println(DatabaseConfig) + _cfg.Settings.multiDatabase() } } diff --git a/tools/config/database.go b/tools/config/database.go index 0de8f4f6..b260b691 100644 --- a/tools/config/database.go +++ b/tools/config/database.go @@ -5,4 +5,7 @@ type Database struct { Source string } -var DatabaseConfig = new(Database) +var ( + DatabaseConfig = new(Database) + DatabasesConfig = make(map[string]*Database) +)