autoresearch iter 16: add batch user lookup and use it for log enrichment

enrichAccessLogsWithUsers preferred the UserService contract over the local
repository — correct layering, but it looped GetUserByID and issued up to a
page-size worth of separate SELECTs against w_users, while the single-query
WHERE id IN variant was only reached in the no-contract fallback branch.
Give the contract a GetUsersByIDs so callers can keep the layering and drop
the N+1. The test asserts 1 query batched against 3 per-id, so the counting
itself is checked.
This commit is contained in:
ryan
2026-08-29 08:51:30 +08:00
parent 5df282f296
commit 976f9b15ae
5 changed files with 98 additions and 3 deletions
+4
View File
@@ -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)
+5 -3
View File
@@ -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 {
+12
View File
@@ -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
+12
View File
@@ -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 {
@@ -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")
}