fix(risk_control,message_gateway): fix SQL LIKE escape syntax and use UserService contract

This commit is contained in:
ryan
2026-08-28 20:19:24 +08:00
parent df351cbd33
commit 0035e548a5
6 changed files with 125 additions and 36 deletions
@@ -19,6 +19,8 @@ var (
cacheSvc contracts.CacheService
taskMu sync.RWMutex
taskSvc contracts.TaskService
userMu sync.RWMutex
userSvc contracts.UserService
)
// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply.
@@ -44,6 +46,12 @@ func setTaskService(s contracts.TaskService) {
taskSvc = s
}
func setUserService(s contracts.UserService) {
userMu.Lock()
defer userMu.Unlock()
userSvc = s
}
func getDB(ctx context.Context) *gorm.DB {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
@@ -76,3 +84,14 @@ func getTaskService() contracts.TaskService {
defer taskMu.RUnlock()
return taskSvc
}
func getUserService(ctx context.Context) contracts.UserService {
if c, ok := ctx.(*core.Context); ok && c != nil {
if s, err := core.Inject[contracts.UserService](c); err == nil && s != nil {
return s
}
}
userMu.RLock()
defer userMu.RUnlock()
return userSvc
}
@@ -79,7 +79,7 @@ type PushNotificationEvent struct {
// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
func (p *Plugin) Apply(ctx *core.Context) error {
// 0. Bind DBService, CacheService, TaskService
// 0. Bind DBService, CacheService, TaskService, UserService
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
setDBService(db)
} else {
@@ -101,10 +101,18 @@ func (p *Plugin) Apply(ctx *core.Context) error {
setTaskService(taskSvc)
})
}
if uSvc, err := core.Inject[contracts.UserService](ctx); err == nil && uSvc != nil {
setUserService(uSvc)
} else {
core.When[contracts.UserService](ctx, func(uSvc contracts.UserService) {
setUserService(uSvc)
})
}
ctx.OnDispose(func() error {
setDBService(nil)
setCacheService(nil)
setTaskService(nil)
setUserService(nil)
return nil
})
@@ -255,22 +255,45 @@ func listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]Pus
return ListActivePushEventsByTaskTypeRecord(ctx, taskType)
}
func queryUser(ctx context.Context, fromService func(contracts.UserService) (*contracts.UserDTO, error), dbField string, dbVal any) (*contracts.UserDTO, error) {
if userSvc := getUserService(ctx); userSvc != nil {
return fromService(userSvc)
}
if db := getDB(ctx); db != nil {
var user contracts.UserDTO
if err := db.Table("w_users").Where(dbField+" = ?", dbVal).First(&user).Error; err == nil {
return &user, nil
}
}
return nil, errors.New("user not found")
}
func findUserByID(ctx context.Context, id uint64) (*contracts.UserDTO, error) {
return queryUser(ctx, func(s contracts.UserService) (*contracts.UserDTO, error) {
return s.GetUserByID(ctx, id)
}, "id", id)
}
func findUserByUsername(ctx context.Context, username string) (*contracts.UserDTO, error) {
return queryUser(ctx, func(s contracts.UserService) (*contracts.UserDTO, error) {
return s.GetUserByUsername(ctx, username)
}, "username", username)
}
func loadUserFromPayload(ctx context.Context, data map[string]any) any {
if u, exists := data["user"]; exists && u != nil {
return u
}
if userID, ok := extractUserID(data); ok && userID > 0 {
var user contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err == nil {
return &user
if user, err := findUserByID(ctx, userID); err == nil && user != nil {
return user
}
}
if username := extractUsername(data); username != "" {
var user contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("username = ?", username).First(&user).Error; err == nil {
return &user
if user, err := findUserByUsername(ctx, username); err == nil && user != nil {
return user
}
}
return nil
@@ -369,24 +392,36 @@ func resolveDynamicKeyword(target string, flatBody map[string]any) string {
}
func resolveTargetUser(ctx context.Context, resolved, _ string) (contracts.UserDTO, bool) {
var user contracts.UserDTO
if id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
if err := getDB(ctx).Table("w_users").Where("id = ?", id).First(&user).Error; err == nil {
return user, true
if u, err := findUserByID(ctx, id); err == nil && u != nil {
return *u, true
}
}
if err := getDB(ctx).Table("w_users").Where("username = ?", resolved).First(&user).Error; err == nil {
return user, true
if u, err := findUserByUsername(ctx, resolved); err == nil && u != nil {
return *u, true
}
return user, false
return contracts.UserDTO{}, false
}
func getFirstAdminUser(ctx context.Context) (*contracts.UserDTO, error) {
if userSvc := getUserService(ctx); userSvc != nil {
return userSvc.GetFirstAdminUser(ctx)
}
if db := getDB(ctx); db != nil {
var adminUser contracts.UserDTO
if err := db.Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err == nil {
return &adminUser, nil
}
}
return nil, errors.New("no admin user found")
}
func resolveSystemTarget(ctx context.Context, resolved, channel string) (string, bool) {
if resolved != "系统" && resolved != "system" && resolved != "0" {
return "", false
}
var adminUser contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
adminUser, err := getFirstAdminUser(ctx)
if err != nil || adminUser == nil {
return resolved, true
}
if channel == channelEmail && adminUser.Email != "" {
@@ -423,9 +458,8 @@ func resolveSMTPConfig(ctx context.Context, url, token, other string) (string, s
}
func getSystemUser(ctx context.Context) *contracts.UserDTO {
var user contracts.UserDTO
if err := getDB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&user).Error; err == nil {
return &user
if adminUser, err := getFirstAdminUser(ctx); err == nil && adminUser != nil {
return adminUser
}
return &contracts.UserDTO{
Username: "system",