mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-08 16:46:37 +08:00
fix(risk_control,message_gateway): fix SQL LIKE escape syntax and use UserService contract
This commit is contained in:
@@ -155,19 +155,41 @@ type accessLogsResponse struct {
|
|||||||
List []accessLogItem `json:"list"`
|
List []accessLogItem `json:"list"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const userQueryMaxLimit = 100
|
||||||
|
|
||||||
|
func findUserIDsByUsername(ctx context.Context, username string) ([]uint64, error) {
|
||||||
|
if userSvc := GetUserService(ctx); userSvc != nil {
|
||||||
|
users, _, err := userSvc.ListUsers(ctx, 1, userQueryMaxLimit, username)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("查询用户信息失败: %w", err)
|
||||||
|
}
|
||||||
|
ids := make([]uint64, 0, len(users))
|
||||||
|
for _, u := range users {
|
||||||
|
ids = append(ids, u.ID)
|
||||||
|
}
|
||||||
|
return ids, nil
|
||||||
|
}
|
||||||
|
gormDB := GetDB(ctx)
|
||||||
|
if gormDB == nil {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
var ids []uint64
|
||||||
|
if err := gormDB.Table("w_users").
|
||||||
|
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
|
||||||
|
Pluck("id", &ids).Error; err != nil {
|
||||||
|
return nil, fmt.Errorf("查询用户信息失败: %w", err)
|
||||||
|
}
|
||||||
|
return ids, nil
|
||||||
|
}
|
||||||
|
|
||||||
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (contracts.AccessLogFilterDTO, error) {
|
func buildAccessLogFilter(ctx context.Context, c *gin.Context) (contracts.AccessLogFilterDTO, error) {
|
||||||
filter := contracts.AccessLogFilterDTO{}
|
filter := contracts.AccessLogFilterDTO{}
|
||||||
|
|
||||||
username := c.Query("username")
|
username := c.Query("username")
|
||||||
if username != "" {
|
if username != "" {
|
||||||
var userIDs []uint64
|
userIDs, err := findUserIDsByUsername(ctx, username)
|
||||||
gormDB := GetDB(ctx)
|
if err != nil {
|
||||||
if gormDB != nil {
|
return filter, err
|
||||||
if err := gormDB.Table("w_users").
|
|
||||||
Where("username LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(username)+"%").
|
|
||||||
Pluck("id", &userIDs).Error; err != nil {
|
|
||||||
return filter, fmt.Errorf("查询用户信息失败: %w", err)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
filter.UserIDs = userIDs
|
filter.UserIDs = userIDs
|
||||||
}
|
}
|
||||||
@@ -214,13 +236,18 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
userMap := make(map[uint64]struct{ Username, Nickname string })
|
userMap := make(map[uint64]struct{ Username, Nickname string })
|
||||||
var users []struct {
|
if userSvc := GetUserService(ctx); userSvc != nil {
|
||||||
ID uint64
|
for _, uid := range userIDs {
|
||||||
Username string
|
if u, err := userSvc.GetUserByID(ctx, uid); err == nil && u != nil {
|
||||||
Nickname string
|
userMap[uid] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
|
||||||
}
|
}
|
||||||
gormDB := GetDB(ctx)
|
}
|
||||||
if gormDB != nil {
|
} else if gormDB := GetDB(ctx); gormDB != nil {
|
||||||
|
var users []struct {
|
||||||
|
ID uint64
|
||||||
|
Username string
|
||||||
|
Nickname string
|
||||||
|
}
|
||||||
if err := gormDB.Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
|
if err := gormDB.Table("w_users").Where("id IN ?", userIDs).Find(&users).Error; err == nil {
|
||||||
for _, u := range users {
|
for _, u := range users {
|
||||||
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
|
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
|
||||||
|
|||||||
@@ -19,6 +19,8 @@ var (
|
|||||||
cacheSvc contracts.CacheService
|
cacheSvc contracts.CacheService
|
||||||
taskMu sync.RWMutex
|
taskMu sync.RWMutex
|
||||||
taskSvc contracts.TaskService
|
taskSvc contracts.TaskService
|
||||||
|
userMu sync.RWMutex
|
||||||
|
userSvc contracts.UserService
|
||||||
)
|
)
|
||||||
|
|
||||||
// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply.
|
// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply.
|
||||||
@@ -44,6 +46,12 @@ func setTaskService(s contracts.TaskService) {
|
|||||||
taskSvc = s
|
taskSvc = s
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func setUserService(s contracts.UserService) {
|
||||||
|
userMu.Lock()
|
||||||
|
defer userMu.Unlock()
|
||||||
|
userSvc = s
|
||||||
|
}
|
||||||
|
|
||||||
func getDB(ctx context.Context) *gorm.DB {
|
func getDB(ctx context.Context) *gorm.DB {
|
||||||
if c, ok := ctx.(*core.Context); ok && c != nil {
|
if c, ok := ctx.(*core.Context); ok && c != nil {
|
||||||
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil {
|
||||||
@@ -76,3 +84,14 @@ func getTaskService() contracts.TaskService {
|
|||||||
defer taskMu.RUnlock()
|
defer taskMu.RUnlock()
|
||||||
return taskSvc
|
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.
|
// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context.
|
||||||
func (p *Plugin) Apply(ctx *core.Context) error {
|
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 {
|
if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil {
|
||||||
setDBService(db)
|
setDBService(db)
|
||||||
} else {
|
} else {
|
||||||
@@ -101,10 +101,18 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
setTaskService(taskSvc)
|
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 {
|
ctx.OnDispose(func() error {
|
||||||
setDBService(nil)
|
setDBService(nil)
|
||||||
setCacheService(nil)
|
setCacheService(nil)
|
||||||
setTaskService(nil)
|
setTaskService(nil)
|
||||||
|
setUserService(nil)
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|||||||
@@ -255,22 +255,45 @@ func listActivePushEventsByTaskType(ctx context.Context, taskType string) ([]Pus
|
|||||||
return ListActivePushEventsByTaskTypeRecord(ctx, taskType)
|
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 {
|
func loadUserFromPayload(ctx context.Context, data map[string]any) any {
|
||||||
if u, exists := data["user"]; exists && u != nil {
|
if u, exists := data["user"]; exists && u != nil {
|
||||||
return u
|
return u
|
||||||
}
|
}
|
||||||
|
|
||||||
if userID, ok := extractUserID(data); ok && userID > 0 {
|
if userID, ok := extractUserID(data); ok && userID > 0 {
|
||||||
var user contracts.UserDTO
|
if user, err := findUserByID(ctx, userID); err == nil && user != nil {
|
||||||
if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err == nil {
|
return user
|
||||||
return &user
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if username := extractUsername(data); username != "" {
|
if username := extractUsername(data); username != "" {
|
||||||
var user contracts.UserDTO
|
if user, err := findUserByUsername(ctx, username); err == nil && user != nil {
|
||||||
if err := getDB(ctx).Table("w_users").Where("username = ?", username).First(&user).Error; err == nil {
|
return user
|
||||||
return &user
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
return nil
|
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) {
|
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 id, err := strconv.ParseUint(resolved, 10, 64); err == nil {
|
||||||
if err := getDB(ctx).Table("w_users").Where("id = ?", id).First(&user).Error; err == nil {
|
if u, err := findUserByID(ctx, id); err == nil && u != nil {
|
||||||
return user, true
|
return *u, true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if err := getDB(ctx).Table("w_users").Where("username = ?", resolved).First(&user).Error; err == nil {
|
if u, err := findUserByUsername(ctx, resolved); err == nil && u != nil {
|
||||||
return user, true
|
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) {
|
func resolveSystemTarget(ctx context.Context, resolved, channel string) (string, bool) {
|
||||||
if resolved != "系统" && resolved != "system" && resolved != "0" {
|
if resolved != "系统" && resolved != "system" && resolved != "0" {
|
||||||
return "", false
|
return "", false
|
||||||
}
|
}
|
||||||
var adminUser contracts.UserDTO
|
adminUser, err := getFirstAdminUser(ctx)
|
||||||
if err := getDB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil {
|
if err != nil || adminUser == nil {
|
||||||
return resolved, true
|
return resolved, true
|
||||||
}
|
}
|
||||||
if channel == channelEmail && adminUser.Email != "" {
|
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 {
|
func getSystemUser(ctx context.Context) *contracts.UserDTO {
|
||||||
var user contracts.UserDTO
|
if adminUser, err := getFirstAdminUser(ctx); err == nil && adminUser != nil {
|
||||||
if err := getDB(ctx).Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&user).Error; err == nil {
|
return adminUser
|
||||||
return &user
|
|
||||||
}
|
}
|
||||||
return &contracts.UserDTO{
|
return &contracts.UserDTO{
|
||||||
Username: "system",
|
Username: "system",
|
||||||
|
|||||||
@@ -93,7 +93,7 @@ func applyFilter(query *gorm.DB, filter AccessLogFilter) *gorm.DB {
|
|||||||
query = query.Where("user_id IN ?", filter.UserIDs)
|
query = query.Where("user_id IN ?", filter.UserIDs)
|
||||||
}
|
}
|
||||||
if filter.Path != "" {
|
if filter.Path != "" {
|
||||||
query = query.Where("path LIKE ?", "%"+util.EscapeLike(filter.Path)+"%")
|
query = query.Where("path LIKE ? ESCAPE '\\'", "%"+util.EscapeLike(filter.Path)+"%")
|
||||||
}
|
}
|
||||||
if filter.StartTime != nil {
|
if filter.StartTime != nil {
|
||||||
query = query.Where("created_at >= ?", *filter.StartTime)
|
query = query.Where("created_at >= ?", *filter.StartTime)
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ package logstore
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"Wavelet/pkg/idgen"
|
"Wavelet/pkg/idgen"
|
||||||
|
"Wavelet/pkg/util"
|
||||||
"context"
|
"context"
|
||||||
"errors"
|
"errors"
|
||||||
"fmt"
|
"fmt"
|
||||||
@@ -153,8 +154,8 @@ func buildUserAccessLogWhere(filter AccessLogFilter) (string, []any, bool) {
|
|||||||
args = append(args, filter.UserIDs)
|
args = append(args, filter.UserIDs)
|
||||||
}
|
}
|
||||||
if trimmed := strings.TrimSpace(filter.Path); trimmed != "" {
|
if trimmed := strings.TrimSpace(filter.Path); trimmed != "" {
|
||||||
parts = append(parts, "path LIKE ?")
|
parts = append(parts, "path LIKE ? ESCAPE '\\'")
|
||||||
args = append(args, "%"+trimmed+"%")
|
args = append(args, "%"+util.EscapeLike(trimmed)+"%")
|
||||||
}
|
}
|
||||||
if filter.StartTime != nil {
|
if filter.StartTime != nil {
|
||||||
parts = append(parts, "created_at >= ?")
|
parts = append(parts, "created_at >= ?")
|
||||||
|
|||||||
Reference in New Issue
Block a user