数据库连接初始化优化

权限部分修改
This commit is contained in:
linwenxiang
2021-02-21 14:06:13 +08:00
parent 4979fd659f
commit deeb82048b
37 changed files with 726 additions and 721 deletions
-167
View File
@@ -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, "", "删除成功")
}
+108 -16
View File
@@ -2,14 +2,13 @@ package sys_role
import ( import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"go-admin/app/admin/models/system" "go-admin/app/admin/models/system"
"go-admin/app/admin/service" "go-admin/app/admin/service"
"go-admin/app/admin/service/dto" "go-admin/app/admin/service/dto"
"go-admin/common/apis" "go-admin/common/apis"
"go-admin/common/global"
"go-admin/common/log" "go-admin/common/log"
"go-admin/tools" "go-admin/tools"
"net/http" "net/http"
) )
@@ -17,6 +16,17 @@ type SysRole struct {
apis.Api 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) { func (e *SysRole) GetSysRoleList(c *gin.Context) {
msgID := tools.GenerateMsgIDFromContext(c) msgID := tools.GenerateMsgIDFromContext(c)
d := new(dto.SysRoleSearch) 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(), "查询成功") 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) { func (e *SysRole) GetSysRole(c *gin.Context) {
control := new(dto.SysRoleById) control := new(dto.SysRoleById)
db, err := tools.GetOrm(c) db, err := tools.GetOrm(c)
@@ -76,6 +94,16 @@ func (e *SysRole) GetSysRole(c *gin.Context) {
e.OK(c, object, "查看成功") 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) { func (e *SysRole) InsertSysRole(c *gin.Context) {
control := new(dto.SysRoleControl) control := new(dto.SysRoleControl)
db, err := tools.GetOrm(c) db, err := tools.GetOrm(c)
@@ -97,21 +125,35 @@ func (e *SysRole) InsertSysRole(c *gin.Context) {
return return
} }
// 设置创建人 // 设置创建人
object.CreateBy= tools.GetUserId(c) object.CreateBy = tools.GetUserId(c)
serviceSysRole := service.SysRole{} s := service.SysRole{}
serviceSysRole.Orm = db s.Orm = db
serviceSysRole.MsgID = msgID s.MsgID = msgID
err = serviceSysRole.InsertSysRole(object) err = s.InsertSysRole(object)
if err != nil { if err != nil {
log.Error(err) log.Error(err)
e.Error(c, http.StatusInternalServerError, err, "创建失败") e.Error(c, http.StatusInternalServerError, err, "创建失败")
return return
} }
_, err = global.LoadPolicy(c)
if err != nil {
e.Error(c, http.StatusInternalServerError, err, "")
return
}
e.OK(c, object.GetId(), "创建成功") 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) { func (e *SysRole) UpdateSysRole(c *gin.Context) {
control := new(dto.SysRoleControl) control := new(dto.SysRoleControl)
db, err := tools.GetOrm(c) db, err := tools.GetOrm(c)
@@ -134,17 +176,30 @@ func (e *SysRole) UpdateSysRole(c *gin.Context) {
} }
object.UpdateBy = tools.GetUserId(c) object.UpdateBy = tools.GetUserId(c)
serviceSysRole := service.SysRole{} s := service.SysRole{}
serviceSysRole.Orm = db s.Orm = db
serviceSysRole.MsgID = msgID s.MsgID = msgID
err = serviceSysRole.UpdateSysRole(object) err = s.UpdateSysRole(object)
if err != nil { if err != nil {
log.Error(err) log.Error(err)
return return
} }
_, err = global.LoadPolicy(c)
if err != nil {
e.Error(c, http.StatusInternalServerError, err, "")
return
}
e.OK(c, object.GetId(), "更新成功") 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) { func (e *SysRole) DeleteSysRole(c *gin.Context) {
control := new(dto.SysRoleById) control := new(dto.SysRoleById)
db, err := tools.GetOrm(c) db, err := tools.GetOrm(c)
@@ -162,14 +217,51 @@ func (e *SysRole) DeleteSysRole(c *gin.Context) {
return return
} }
serviceSysRole := service.SysRole{} s := service.SysRole{}
serviceSysRole.Orm = db s.Orm = db
serviceSysRole.MsgID = msgID s.MsgID = msgID
err = serviceSysRole.RemoveSysRole(control) err = s.RemoveSysRole(control)
if err != nil { if err != nil {
log.Error(err) log.Error(err)
return return
} }
_, err = global.LoadPolicy(c)
if err != nil {
e.Error(c, http.StatusInternalServerError, err, "")
return
}
e.OK(c, control.GetId(), "删除成功") 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, "操作成功")
}
-42
View File
@@ -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
}
-23
View File
@@ -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")
}
}
-28
View File
@@ -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")
}
}
+3 -1
View File
@@ -17,5 +17,7 @@ func InitMiddleware(r *gin.Engine) {
// Secure is a middleware function that appends security // Secure is a middleware function that appends security
r.Use(Secure) r.Use(Secure)
// 链路追踪 // 链路追踪
r.Use(middleware.Trace()) //r.Use(middleware.Trace())
// 数据库链接
r.Use(middleware.WithContextDb)
} }
+2 -2
View File
@@ -17,7 +17,7 @@ func AuthCheckRole() gin.HandlerFunc {
return func(c *gin.Context) { return func(c *gin.Context) {
data, _ := c.Get(jwtauth.JwtPayloadKey) data, _ := c.Get(jwtauth.JwtPayloadKey)
v := data.(jwtauth.MapClaims) v := data.(jwtauth.MapClaims)
e := global.CasbinEnforcer e := global.Cfg.GetCasbinKey(c.Request.Host)
var res bool var res bool
var err error var err error
msgID := tools.GenerateMsgIDFromContext(c) 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) 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() c.Next()
} else { } 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{ c.JSON(http.StatusOK, gin.H{
"code": 403, "code": 403,
"msg": "对不起,您没有该接口访问权限,请联系管理员", "msg": "对不起,您没有该接口访问权限,请联系管理员",
-47
View File
@@ -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
}
+11
View File
@@ -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"
}
+172
View File
@@ -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
}
-1
View File
@@ -32,7 +32,6 @@ func InitRouter() {
if config.SslConfig.Enable { if config.SslConfig.Enable {
r.Use(handler.TlsHandler()) r.Use(handler.TlsHandler())
} }
r.Use(middleware.WithContextDb(middleware.GetGormFromConfig(global.Cfg)))
r.Use(common.Sentinel()) r.Use(common.Sentinel())
middleware.InitMiddleware(r) middleware.InitMiddleware(r)
+1
View File
@@ -22,4 +22,5 @@ func registerSysRoleRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddle
r.PUT("/:id", api.UpdateSysRole) r.PUT("/:id", api.UpdateSysRole)
r.DELETE("/:id", api.DeleteSysRole) r.DELETE("/:id", api.DeleteSysRole)
} }
v1.PUT("/roledatascope", api.UpdateRoleDataScope)
} }
+10 -10
View File
@@ -117,7 +117,7 @@ func registerBaseRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddlewar
{ {
v1auth.GET("/getinfo", system.GetInfo) v1auth.GET("/getinfo", system.GetInfo)
v1auth.PUT("/roledatascope", system.UpdateRoleDataScope) //v1auth.PUT("/roledatascope", system.UpdateRoleDataScope)
v1auth.GET("/roleMenuTreeselect/:roleId", system.GetMenuTreeRoleselect) v1auth.GET("/roleMenuTreeselect/:roleId", system.GetMenuTreeRoleselect)
v1auth.GET("/roleDeptTreeselect/:roleId", system.GetDeptTreeRoleselect) 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) { //func registerRoleRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddleware) {
role := v1.Group("/role").Use(authMiddleware.MiddlewareFunc()).Use(middleware.AuthCheckRole()) // role := v1.Group("/role").Use(authMiddleware.MiddlewareFunc()).Use(middleware.AuthCheckRole())
{ // {
role.GET("/:roleId", system.GetRole) // role.GET("/:roleId", system.GetRole)
role.POST("", system.InsertRole) // role.POST("", system.InsertRole)
role.PUT("", system.UpdateRole) // role.PUT("", system.UpdateRole)
role.DELETE("/:roleId", system.DeleteRole) // role.DELETE("/:roleId", system.DeleteRole)
} // }
} //}
func registerSysUserRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddleware) { func registerSysUserRouter(v1 *gin.RouterGroup, authMiddleware *jwt.GinJWTMiddleware) {
sysuser := v1.Group("/sysUser").Use(authMiddleware.MiddlewareFunc()).Use(middleware.AuthCheckRole()) sysuser := v1.Group("/sysUser").Use(authMiddleware.MiddlewareFunc()).Use(middleware.AuthCheckRole())
+3
View File
@@ -51,6 +51,7 @@ type SysRoleControl struct {
Remark string `form:"remark" comment:"备注"` // 备注 Remark string `form:"remark" comment:"备注"` // 备注
Admin bool `form:"admin" comment:"是否管理员"` Admin bool `form:"admin" comment:"是否管理员"`
DataScope string `form:"dataScope" comment:"是否管理员"` DataScope string `form:"dataScope" comment:"是否管理员"`
MenuIds []int `json:"menuIds"`
} }
// Bind 映射上下文中的结构体数据 // Bind 映射上下文中的结构体数据
@@ -86,6 +87,7 @@ func (s *SysRoleControl) Generate() (*system.SysRole, error) {
Remark: s.Remark, Remark: s.Remark,
Admin: s.Admin, Admin: s.Admin,
DataScope: s.DataScope, DataScope: s.DataScope,
MenuIds: s.MenuIds,
}, nil }, nil
} }
@@ -131,4 +133,5 @@ func (s *SysRoleById) GenerateM() (*models.SysRole, error) {
type RoleDataScopeReq struct { type RoleDataScopeReq struct {
RoleId int `json:"roleId" binding:"required"` RoleId int `json:"roleId" binding:"required"`
DataScope string `json:"dataScope" binding:"required"` DataScope string `json:"dataScope" binding:"required"`
DeptIds []int `json:"deptIds"`
} }
+131 -8
View File
@@ -2,12 +2,16 @@ package service
import ( import (
"errors" "errors"
"gorm.io/gorm"
"go-admin/app/admin/models"
"go-admin/app/admin/models/system" "go-admin/app/admin/models/system"
"go-admin/app/admin/service/dto" "go-admin/app/admin/service/dto"
cDto "go-admin/common/dto" cDto "go-admin/common/dto"
orm "go-admin/common/global"
"go-admin/common/log" "go-admin/common/log"
"go-admin/common/service" "go-admin/common/service"
"gorm.io/gorm"
) )
type SysRole struct { 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) log.Errorf("msgID[%s] db error:%s", msgID, err)
return err return err
} }
if db.Error != nil { if err != nil {
log.Errorf("msgID[%s] db error:%s", msgID, err) log.Errorf("msgID[%s] db error:%s", msgID, err)
return 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 return nil
} }
// InsertSysRole 创建SysRole对象 // InsertSysRole 创建SysRole对象
func (e *SysRole) InsertSysRole(model *system.SysRole) error { func (e *SysRole) InsertSysRole(c *system.SysRole) error {
var err error var err error
var data system.SysRole var data system.SysRole
msgID := e.MsgID msgID := e.MsgID
err = e.Orm.Model(&data). tx := e.Orm.Begin()
Create(model).Error defer func() {
if err != nil {
tx.Rollback()
} else {
tx.Commit()
}
}()
err = tx.Model(&data).
Create(c).Error
if err != nil { if err != nil {
log.Errorf("msgID[%s] db error:%s", msgID, err) log.Errorf("msgID[%s] db error:%s", msgID, err)
return 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 return nil
} }
// UpdateSysRole 修改SysRole对象 // UpdateSysRole 修改SysRole对象
func (e *SysRole) UpdateSysRole(c *system.SysRole) error { func (e *SysRole) UpdateSysRole(c *system.SysRole) error {
var err error var err error
var data system.SysRole
msgID := e.MsgID 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) Where(c.GetId()).Updates(c)
if db.Error != nil { if db.Error != nil {
log.Errorf("msgID[%s] db error:%s", msgID, err) 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 { if db.RowsAffected == 0 {
return errors.New("无权更新该数据") 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 return nil
} }
@@ -94,7 +146,16 @@ func (e *SysRole) RemoveSysRole(d *dto.SysRoleById) error {
var data system.SysRole var data system.SysRole
msgID := e.MsgID 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 { if db.Error != nil {
err = db.Error err = db.Error
log.Errorf("MsgID[%s] Delete error: %s", msgID, err) log.Errorf("MsgID[%s] Delete error: %s", msgID, err)
@@ -104,5 +165,67 @@ func (e *SysRole) RemoveSysRole(d *dto.SysRoleById) error {
err = errors.New("无权删除该数据") err = errors.New("无权删除该数据")
return err 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 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
}
+113
View File
@@ -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
}
+4 -20
View File
@@ -17,19 +17,14 @@ import (
"go-admin/common/database" "go-admin/common/database"
"go-admin/common/global" "go-admin/common/global"
"go-admin/common/log" "go-admin/common/log"
mycasbin "go-admin/pkg/casbin"
"go-admin/pkg/logger" "go-admin/pkg/logger"
"go-admin/tools" "go-admin/tools"
"go-admin/tools/config" "go-admin/tools/config"
"go-admin/tools/trace"
) )
var ( var (
configYml string configYml string
port string StartCmd = &cobra.Command{
mode string
traceStart bool
StartCmd = &cobra.Command{
Use: "server", Use: "server",
Short: "Start API server", Short: "Start API server",
Example: "go-admin server -c config/settings.yml", Example: "go-admin server -c config/settings.yml",
@@ -47,9 +42,6 @@ var AppRouters = make([]func(), 0)
func init() { func init() {
StartCmd.PersistentFlags().StringVarP(&configYml, "config", "c", "config/settings.yml", "Start server with provided configuration file") 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方法 //注册路由 fixme 其他应用的路由,在本目录新建文件放在init方法
AppRouters = append(AppRouters, router.InitRouter) AppRouters = append(AppRouters, router.InitRouter)
@@ -65,9 +57,7 @@ func setup() {
global.JobLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "job") global.JobLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "job")
global.RequestLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "request") global.RequestLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "request")
//3. 初始化数据库链接 //3. 初始化数据库链接
database.Setup(config.DatabaseConfig.Driver) database.Setup()
//4. 接口访问控制加载
global.CasbinEnforcer = mycasbin.Setup(global.Eloquent, "sys_")
usageStr := `starting api server` usageStr := `starting api server`
log.Info(usageStr) log.Info(usageStr)
@@ -85,7 +75,7 @@ func run() error {
engine = gin.New() engine = gin.New()
} }
if mode == "dev" { if config.ApplicationConfig.Mode == "dev" {
//监控 //监控
AppRouters = append(AppRouters, router.Monitor) AppRouters = append(AppRouters, router.Monitor)
} }
@@ -107,12 +97,6 @@ func run() error {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel() defer cancel()
if traceStart {
//链路追踪, fixme 页面显示需要自备梯子
trace.Start()
defer trace.Stop(ctx)
}
go func() { go func() {
// 服务连接 // 服务连接
if config.SslConfig.Enable { if config.SslConfig.Enable {
+2 -2
View File
@@ -12,7 +12,6 @@ import (
var ( var (
configYml string configYml string
mode string
StartCmd = &cobra.Command{ StartCmd = &cobra.Command{
Use: "config", Use: "config",
Short: "Get Application config info", Short: "Get Application config info",
@@ -42,7 +41,8 @@ func run() {
} }
fmt.Println("jwt:", string(jwt)) 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 { if errs != nil {
fmt.Println(errs.Error()) fmt.Println(errs.Error())
} }
@@ -28,7 +28,7 @@ func _1599190683659Tables(db *gorm.DB, version string) error {
new(system.SysLoginLog), new(system.SysLoginLog),
new(system.SysOperaLog), new(system.SysOperaLog),
new(models.RoleMenu), new(models.RoleMenu),
new(models.SysRoleDept), new(system.SysRoleDept),
new(models.SysUser), new(models.SysUser),
new(system.SysRole), new(system.SysRole),
new(models.Post), new(models.Post),
+1 -1
View File
@@ -72,7 +72,7 @@ func migrateModel() error {
} }
func initDB() error { func initDB() error {
//3. 初始化数据库链接 //3. 初始化数据库链接
database.Setup(config.DatabaseConfig.Driver) database.Setup()
//4. 数据库迁移 //4. 数据库迁移
fmt.Println("数据库迁移开始") fmt.Println("数据库迁移开始")
_ = migrateModel() _ = migrateModel()
+27 -33
View File
@@ -1,47 +1,47 @@
package config package config
import ( import (
"database/sql"
"net/http" "net/http"
"github.com/casbin/casbin/v2"
"github.com/go-admin-team/go-admin-core/logger" "github.com/go-admin-team/go-admin-core/logger"
"gorm.io/gorm"
) )
type Config struct { type Config struct {
saas bool dbs map[string]*gorm.DB
dbs map[string]*DBConfig casbins map[string]*casbin.SyncedEnforcer
db *DBConfig engine http.Handler
engine http.Handler
} }
type DBConfig struct { // SetDb 设置对应key的db
Driver string func (c *Config) SetDb(key string, db *gorm.DB) {
DB *sql.DB
}
// SetDbs 设置对应key的db
func (c *Config) SetDbs(key string, db *DBConfig) {
c.dbs[key] = db c.dbs[key] = db
} }
// GetDbs 获取所有map里的db数据 // GetDb 获取所有map里的db数据
func (c *Config) GetDbs() map[string]*DBConfig { func (c *Config) GetDb() map[string]*gorm.DB {
return c.dbs return c.dbs
} }
// GetDbByKey 根据key获取db // 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] return c.dbs[key]
} }
// SetDb 设置单个db func (c *Config) SetCasbin(key string, enforcer *casbin.SyncedEnforcer) {
func (c *Config) SetDb(db *DBConfig) { c.casbins[key] = enforcer
c.db = db
} }
// GetDb 获取单个db // GetCasbinKey 根据key获取casbin
func (c *Config) GetDb() *DBConfig { func (c *Config) GetCasbinKey(key string) *casbin.SyncedEnforcer {
return c.db if e, ok := c.casbins["*"]; ok {
return e
}
return c.casbins[key]
} }
// SetEngine 设置路由引擎 // SetEngine 设置路由引擎
@@ -64,16 +64,10 @@ func (c *Config) GetLogger() logger.Logger {
return logger.DefaultLogger return logger.DefaultLogger
} }
// SetSaas 设置是否是saas应用 // NewConfig 默认值
func (c *Config) SetSaas(saas bool) { func NewConfig() *Config {
c.saas = saas return &Config{
} dbs: make(map[string]*gorm.DB),
casbins: make(map[string]*casbin.SyncedEnforcer),
// GetSaas 获取是否是saas应用 }
func (c *Config) GetSaas() bool {
return c.saas
}
func DefaultConfig() *Config {
return &Config{}
} }
+7 -8
View File
@@ -1,22 +1,21 @@
package config package config
import ( import (
"github.com/casbin/casbin/v2"
"net/http" "net/http"
"github.com/go-admin-team/go-admin-core/logger" "github.com/go-admin-team/go-admin-core/logger"
"gorm.io/gorm"
) )
type Conf interface { type Conf interface {
//多db设置,⚠️SetDbs不允许并发,可以根据自己的业务,例如app分库、host分库 //多db设置,⚠️SetDbs不允许并发,可以根据自己的业务,例如app分库、host分库
SetDbs(key string, db *DBConfig) SetDb(key string, db *gorm.DB)
GetDbs() map[string]*DBConfig GetDb() map[string]*gorm.DB
GetDbByKey(key string) *DBConfig GetDbByKey(key string) *gorm.DB
GetSaas() bool
SetSaas(bool)
//单库业务实现这两个接口 SetCasbin(key string, enforcer *casbin.SyncedEnforcer)
SetDb(db *DBConfig) GetCasbinKey(key string) *casbin.SyncedEnforcer
GetDb() *DBConfig
//使用的路由 //使用的路由
SetEngine(engine http.Handler) SetEngine(engine http.Handler)
+54 -12
View File
@@ -1,17 +1,59 @@
// +build !sqlite3
package database package database
// Setup 配置数据库 import (
func Setup(driver string) { . "log"
dbType := driver "time"
if dbType == "mysql" {
var db = new(Mysql)
db.Setup()
}
if dbType == "postgres" { logCore "github.com/go-admin-team/go-admin-core/logger"
var db = new(PgSql) "gorm.io/gorm"
db.Setup() "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)
}
-21
View File
@@ -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()
}
}
-11
View File
@@ -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
}
-77
View File
@@ -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
}
+14
View File
@@ -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,
}
+16
View File
@@ -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,
}
-75
View File
@@ -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
}
-77
View File
@@ -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
}
+2 -3
View File
@@ -2,9 +2,8 @@ package dto
import ( import (
"github.com/go-admin-team/go-admin-core/tools/search" "github.com/go-admin-team/go-admin-core/tools/search"
"go-admin/common/global"
"gorm.io/gorm" "gorm.io/gorm"
"go-admin/tools/config"
) )
type GeneralDelDto struct { type GeneralDelDto struct {
@@ -45,7 +44,7 @@ func MakeCondition(q interface{}) func(db *gorm.DB) *gorm.DB {
GormPublic: search.GormPublic{}, GormPublic: search.GormPublic{},
Join: make([]*search.GormJoin, 0), Join: make([]*search.GormJoin, 0),
} }
search.ResolveSearchQuery(config.DatabaseConfig.Driver, q, condition) search.ResolveSearchQuery(global.Driver, q, condition)
for _, join := range condition.Join { for _, join := range condition.Join {
if join == nil { if join == nil {
continue continue
+1 -3
View File
@@ -1,7 +1,6 @@
package global package global
import ( import (
"github.com/casbin/casbin/v2"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/robfig/cron/v3" "github.com/robfig/cron/v3"
"gorm.io/gorm" "gorm.io/gorm"
@@ -15,10 +14,9 @@ const (
Version = "1.2.3" Version = "1.2.3"
) )
var Cfg config.Conf = config.DefaultConfig() var Cfg config.Conf = config.NewConfig()
var GinEngine *gin.Engine var GinEngine *gin.Engine
var CasbinEnforcer *casbin.SyncedEnforcer
var Eloquent *gorm.DB var Eloquent *gorm.DB
var GADMCron *cron.Cron var GADMCron *cron.Cron
+8 -4
View File
@@ -3,13 +3,17 @@ package global
import ( import (
"github.com/casbin/casbin/v2" "github.com/casbin/casbin/v2"
"github.com/casbin/casbin/v2/log" "github.com/casbin/casbin/v2/log"
"github.com/gin-gonic/gin"
"go-admin/tools"
) )
func LoadPolicy() (*casbin.SyncedEnforcer, error) { func LoadPolicy(c *gin.Context) (*casbin.SyncedEnforcer, error) {
if err := CasbinEnforcer.LoadPolicy(); err == nil { if err := Cfg.GetCasbinKey(c.Request.Host).LoadPolicy(); err == nil {
return CasbinEnforcer, err return Cfg.GetCasbinKey(c.Request.Host), err
} else { } 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 return nil, err
} }
} }
+5 -10
View File
@@ -2,16 +2,11 @@ package middleware
import ( import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"gorm.io/gorm"
"go-admin/common/global"
) )
func WithContextDb(dbMap map[string]*gorm.DB) gin.HandlerFunc { func WithContextDb(c *gin.Context) {
return func(c *gin.Context) { c.Set("db", global.Cfg.GetDbByKey(c.Request.Host))
if db, ok := dbMap["*"]; ok { c.Next()
c.Set("db", db)
} else {
c.Set("db", dbMap[c.Request.Host])
}
c.Next()
}
} }
+2 -3
View File
@@ -16,12 +16,11 @@ import (
) )
func main() { func main() {
var err error db, err := gorm.Open(mysql.Open("root:123456@tcp/inmg?charset=utf8&parseTime=True&loc=Local"), &gorm.Config{})
global.Eloquent, err = gorm.Open(mysql.Open("root:123456@tcp/inmg?charset=utf8&parseTime=True&loc=Local"), &gorm.Config{})
if err != nil { if err != nil {
panic(err) panic(err)
} }
global.CasbinEnforcer = mycasbin.Setup(global.Eloquent, "sys_") _ = mycasbin.Setup(db, "sys_")
global.Logger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "bus") global.Logger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "bus")
global.JobLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "job") global.JobLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "job")
global.RequestLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "request") global.RequestLogger.Logger = logger.SetupLogger(config.LoggerConfig.Path, "request")
+24 -14
View File
@@ -9,6 +9,8 @@ import (
var ( var (
ExtendConfig interface{} ExtendConfig interface{}
_watch config.Watcher
_cfg *Settings
) )
// Settings 兼容原先的配置结构 // Settings 兼容原先的配置结构
@@ -18,19 +20,25 @@ type Settings struct {
// Config 配置集合 // Config 配置集合
type Config struct { type Config struct {
Application *Application `yaml:"application"` Application *Application `yaml:"application"`
Ssl *Ssl `yaml:"ssl"` Ssl *Ssl `yaml:"ssl"`
Logger *Logger `yaml:"logger"` Logger *Logger `yaml:"logger"`
Jwt *Jwt `yaml:"jwt"` Jwt *Jwt `yaml:"jwt"`
Database *Database `yaml:"database"` Database *Database `yaml:"database"`
Gen *Gen `yaml:"gen"` Databases *map[string]*Database `yaml:"databases"`
Extend interface{} `yaml:"extend"` Gen *Gen `yaml:"gen"`
Extend interface{} `yaml:"extend"`
} }
var ( // 多db改造
_watch config.Watcher func (e *Config) multiDatabase() {
_cfg *Settings if len(*e.Databases) == 0 {
) *e.Databases = map[string]*Database{
"*": e.Database,
}
}
}
// Setup 载入配置文件 // Setup 载入配置文件
func Setup(f func(opts ...source.Option) source.Source, options ...source.Option) { 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, Ssl: SslConfig,
Logger: LoggerConfig, Logger: LoggerConfig,
Jwt: JwtConfig, Jwt: JwtConfig,
Database: DatabaseConfig, Database: new(Database),
Databases: &DatabasesConfig,
Gen: GenConfig, Gen: GenConfig,
Extend: ExtendConfig, Extend: ExtendConfig,
}} }}
@@ -58,6 +67,7 @@ func Setup(f func(opts ...source.Option) source.Source, options ...source.Option
if err != nil { if err != nil {
log.Fatal(fmt.Sprintf("Scan config fail: %s", err.Error())) log.Fatal(fmt.Sprintf("Scan config fail: %s", err.Error()))
} }
_cfg.Settings.multiDatabase()
_watch, err = c.Watch() _watch, err = c.Watch()
if err != nil { 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() { func Watch() {
for { for {
v, err := _watch.Next() v, err := _watch.Next()
@@ -80,7 +90,7 @@ func Watch() {
log.Println(fmt.Sprintf("Scan config fail: %s", err.Error())) log.Println(fmt.Sprintf("Scan config fail: %s", err.Error()))
break break
} }
fmt.Println(DatabaseConfig) _cfg.Settings.multiDatabase()
} }
} }
+4 -1
View File
@@ -5,4 +5,7 @@ type Database struct {
Source string Source string
} }
var DatabaseConfig = new(Database) var (
DatabaseConfig = new(Database)
DatabasesConfig = make(map[string]*Database)
)