diff --git a/backend/plugins/domain/admin/handlers_logs.go b/backend/plugins/domain/admin/handlers_logs.go index c4a7a551..a0281e2f 100644 --- a/backend/plugins/domain/admin/handlers_logs.go +++ b/backend/plugins/domain/admin/handlers_logs.go @@ -155,19 +155,41 @@ type accessLogsResponse struct { 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) { filter := contracts.AccessLogFilterDTO{} username := c.Query("username") if username != "" { - var userIDs []uint64 - gormDB := GetDB(ctx) - if gormDB != nil { - 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) - } + userIDs, err := findUserIDsByUsername(ctx, username) + if err != nil { + return filter, err } filter.UserIDs = userIDs } @@ -214,13 +236,18 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) { } userMap := make(map[uint64]struct{ Username, Nickname string }) - var users []struct { - ID uint64 - Username string - Nickname string - } - gormDB := GetDB(ctx) - if gormDB != nil { + if userSvc := GetUserService(ctx); userSvc != nil { + for _, uid := range userIDs { + if u, err := userSvc.GetUserByID(ctx, uid); err == nil && u != nil { + userMap[uid] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname} + } + } + } 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 { for _, u := range users { userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname} diff --git a/backend/plugins/domain/message_gateway/db_helper.go b/backend/plugins/domain/message_gateway/db_helper.go index eb7cef84..a1602914 100644 --- a/backend/plugins/domain/message_gateway/db_helper.go +++ b/backend/plugins/domain/message_gateway/db_helper.go @@ -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 +} diff --git a/backend/plugins/domain/message_gateway/plugin.go b/backend/plugins/domain/message_gateway/plugin.go index 2da67574..e9caee2b 100644 --- a/backend/plugins/domain/message_gateway/plugin.go +++ b/backend/plugins/domain/message_gateway/plugin.go @@ -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 }) diff --git a/backend/plugins/domain/message_gateway/push_logics.go b/backend/plugins/domain/message_gateway/push_logics.go index 688019f3..8271a5cf 100644 --- a/backend/plugins/domain/message_gateway/push_logics.go +++ b/backend/plugins/domain/message_gateway/push_logics.go @@ -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", diff --git a/backend/plugins/domain/risk_control/logstore/access_log.go b/backend/plugins/domain/risk_control/logstore/access_log.go index 154f71d8..2d23be31 100644 --- a/backend/plugins/domain/risk_control/logstore/access_log.go +++ b/backend/plugins/domain/risk_control/logstore/access_log.go @@ -93,7 +93,7 @@ func applyFilter(query *gorm.DB, filter AccessLogFilter) *gorm.DB { query = query.Where("user_id IN ?", filter.UserIDs) } 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 { query = query.Where("created_at >= ?", *filter.StartTime) diff --git a/backend/plugins/domain/risk_control/logstore/gorm.go b/backend/plugins/domain/risk_control/logstore/gorm.go index cf79ec6c..a6ed8990 100644 --- a/backend/plugins/domain/risk_control/logstore/gorm.go +++ b/backend/plugins/domain/risk_control/logstore/gorm.go @@ -5,6 +5,7 @@ package logstore import ( "Wavelet/pkg/idgen" + "Wavelet/pkg/util" "context" "errors" "fmt" @@ -153,8 +154,8 @@ func buildUserAccessLogWhere(filter AccessLogFilter) (string, []any, bool) { args = append(args, filter.UserIDs) } if trimmed := strings.TrimSpace(filter.Path); trimmed != "" { - parts = append(parts, "path LIKE ?") - args = append(args, "%"+trimmed+"%") + parts = append(parts, "path LIKE ? ESCAPE '\\'") + args = append(args, "%"+util.EscapeLike(trimmed)+"%") } if filter.StartTime != nil { parts = append(parts, "created_at >= ?")