diff --git a/backend/core/contracts/user.go b/backend/core/contracts/user.go index 701ee46c..f3fbb51a 100644 --- a/backend/core/contracts/user.go +++ b/backend/core/contracts/user.go @@ -62,6 +62,10 @@ type UserService interface { // GetUserByID retrieves a user by ID. GetUserByID(ctx context.Context, id uint64) (*UserDTO, error) + // GetUsersByIDs retrieves several users in one round-trip. An empty ids + // slice yields no results and touches no storage. + GetUsersByIDs(ctx context.Context, ids []uint64) ([]*UserDTO, error) + // GetUserByUsername retrieves a user by username. GetUserByUsername(ctx context.Context, username string) (*UserDTO, error) diff --git a/backend/plugins/domain/admin/service/log.go b/backend/plugins/domain/admin/service/log.go index 09128668..82a6dd84 100644 --- a/backend/plugins/domain/admin/service/log.go +++ b/backend/plugins/domain/admin/service/log.go @@ -220,9 +220,11 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []model.AccessLogItem) userMap := make(map[uint64]repository.UserDisplayName, len(userIDs)) if userSvc := GetUserService(ctx); userSvc != nil { - for _, uid := range userIDs { - if u, err := userSvc.GetUserByID(ctx, uid); err == nil && u != nil { - userMap[uid] = repository.UserDisplayName{Username: u.Username, Nickname: u.Nickname} + if users, err := userSvc.GetUsersByIDs(ctx, userIDs); err == nil { + for _, u := range users { + if u != nil { + userMap[u.ID] = repository.UserDisplayName{Username: u.Username, Nickname: u.Nickname} + } } } } else if names, err := repository.LoadUserDisplayNames(ctx, userIDs); err == nil { diff --git a/backend/plugins/domain/user/repository.go b/backend/plugins/domain/user/repository.go index c77d5461..2bb02a9e 100644 --- a/backend/plugins/domain/user/repository.go +++ b/backend/plugins/domain/user/repository.go @@ -52,6 +52,18 @@ func GetUserByID(ctx context.Context, id uint64) (*User, error) { return &u, nil } +// GetUsersByIDs 一次性批量获取多个用户,避免调用方按 ID 逐条查询。 +func GetUsersByIDs(ctx context.Context, ids []uint64) ([]User, error) { + if len(ids) == 0 { + return []User{}, nil + } + var users []User + if err := getDB(ctx).Where("id IN ?", ids).Find(&users).Error; err != nil { + return nil, err + } + return users, nil +} + // GetUserByUsername 通过用户名获取用户 func GetUserByUsername(ctx context.Context, username string) (*User, error) { var u User diff --git a/backend/plugins/domain/user/service.go b/backend/plugins/domain/user/service.go index 07281d95..d3b2fd15 100644 --- a/backend/plugins/domain/user/service.go +++ b/backend/plugins/domain/user/service.go @@ -62,6 +62,18 @@ func (s *userServiceImpl) GetUserByID(ctx context.Context, id uint64) (*contract return toUserDTO(u), nil } +func (s *userServiceImpl) GetUsersByIDs(ctx context.Context, ids []uint64) ([]*contracts.UserDTO, error) { + users, err := GetUsersByIDs(ctx, ids) + if err != nil { + return nil, err + } + dtos := make([]*contracts.UserDTO, 0, len(users)) + for i := range users { + dtos = append(dtos, toUserDTO(&users[i])) + } + return dtos, nil +} + func (s *userServiceImpl) GetUserByUsername(ctx context.Context, username string) (*contracts.UserDTO, error) { u, err := GetUserByUsername(ctx, username) if err != nil { diff --git a/backend/plugins/domain/user/users_by_ids_test.go b/backend/plugins/domain/user/users_by_ids_test.go new file mode 100644 index 00000000..9a96e1fa --- /dev/null +++ b/backend/plugins/domain/user/users_by_ids_test.go @@ -0,0 +1,65 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user_test + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "Wavelet/plugins/domain/user" + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" + + database "Wavelet/plugins/infra/database" +) + +// TestGetUsersByIDsUsesSingleQuery 回归:批量取用户必须只发一条 SQL, +// 否则调用方(如访问日志按用户补全)会按 ID 逐条打库。 +func TestGetUsersByIDsUsesSingleQuery(t *testing.T) { + ctx := core.NewContext(context.Background()) + testDB := setupTestDB(t) + require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx)) + require.NoError(t, user.New().Apply(ctx)) + + userSvc, err := core.Inject[contracts.UserService](ctx) + require.NoError(t, err) + + bg := context.Background() + ids := make([]uint64, 0, 3) + for _, name := range []string{"batch_a", "batch_b", "batch_c"} { + u, err := userSvc.CreateUser(bg, contracts.CreateUserRequest{ + Username: name, + Password: "Password789!", + Email: name + "@example.com", + }) + require.NoError(t, err) + ids = append(ids, u.ID) + } + + queries := 0 + require.NoError(t, testDB.Callback().Query().Register("count_queries", func(_ *gorm.DB) { + queries++ + })) + + got, err := userSvc.GetUsersByIDs(bg, ids) + require.NoError(t, err) + assert.Len(t, got, 3) + assert.Equal(t, 1, queries, "batch lookup must issue exactly one query") + + queries = 0 + for _, id := range ids { + _, err := userSvc.GetUserByID(bg, id) + require.NoError(t, err) + } + assert.Equal(t, 3, queries, "per-id lookups cost one query each") + + queries = 0 + empty, err := userSvc.GetUsersByIDs(bg, nil) + require.NoError(t, err) + assert.Empty(t, empty) + assert.Equal(t, 0, queries, "empty id list must not touch storage") +}