From 9ea0e2bdfff79fb0dedc5b023981b482e2fb4653 Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 29 Aug 2026 18:49:37 +0800 Subject: [PATCH] autoresearch iter 30: enforce user lookup column allow-list instead of trusting a comment --- .../domain/message_gateway/errs/errs.go | 4 ++ .../domain/message_gateway/repository/push.go | 13 +++- .../message_gateway/repository/push_test.go | 72 +++++++++++++++++++ 3 files changed, 88 insertions(+), 1 deletion(-) create mode 100644 backend/plugins/domain/message_gateway/repository/push_test.go diff --git a/backend/plugins/domain/message_gateway/errs/errs.go b/backend/plugins/domain/message_gateway/errs/errs.go index 59da3cf8..e6fa4280 100644 --- a/backend/plugins/domain/message_gateway/errs/errs.go +++ b/backend/plugins/domain/message_gateway/errs/errs.go @@ -20,6 +20,10 @@ var ( // ErrRecordNotFound maps GORM's missing-row sentinel at the repository boundary so // upper layers never import gorm. Its text matches gorm.ErrRecordNotFound verbatim. ErrRecordNotFound = errors.New("record not found") + + // ErrUnsupportedUserLookupField rejects a column name that the repository is not + // allowed to interpolate into a WHERE clause. + ErrUnsupportedUserLookupField = errors.New("unsupported user lookup field") ) // User-facing validation and error message constants. diff --git a/backend/plugins/domain/message_gateway/repository/push.go b/backend/plugins/domain/message_gateway/repository/push.go index 187246e7..4e7789f5 100644 --- a/backend/plugins/domain/message_gateway/repository/push.go +++ b/backend/plugins/domain/message_gateway/repository/push.go @@ -295,9 +295,20 @@ func LoadSMTPConfigRecord(ctx context.Context) model.SMTPConfig { return cfg } +// userLookupColumns allow-lists the columns FindUserByFieldRecord may filter on. +// The column name is concatenated into the WHERE clause, so anything not listed +// here must never reach the database. +var userLookupColumns = map[string]struct{}{ + "id": {}, + "username": {}, +} + // FindUserByFieldRecord is the user lookup fallback for when the UserService -// contract is not wired yet. field comes from call sites, never from user input. +// contract is not wired yet. field must be one of userLookupColumns. func FindUserByFieldRecord(ctx context.Context, field string, value any) (*contracts.UserDTO, error) { + if _, ok := userLookupColumns[field]; !ok { + return nil, errs.ErrUnsupportedUserLookupField + } db := GetDB(ctx) if db == nil { return nil, errs.ErrRecordNotFound diff --git a/backend/plugins/domain/message_gateway/repository/push_test.go b/backend/plugins/domain/message_gateway/repository/push_test.go new file mode 100644 index 00000000..7666ddc6 --- /dev/null +++ b/backend/plugins/domain/message_gateway/repository/push_test.go @@ -0,0 +1,72 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository_test + +import ( + "Wavelet/pkg/testhelper" + "Wavelet/plugins/domain/message_gateway/errs" + "Wavelet/plugins/domain/message_gateway/repository" + "context" + "errors" + "testing" + + "gorm.io/gorm" +) + +// stubDBService satisfies contracts.DBService over a test database handle. +type stubDBService struct{ db *gorm.DB } + +func (s stubDBService) GORM() *gorm.DB { return s.db } + +func (s stubDBService) DB(_ context.Context) *gorm.DB { return s.db } + +func (s stubDBService) Named(_ string) *gorm.DB { return s.db } + +// TestFindUserByFieldRecordRejectsUnlistedColumns pins the column allow-list. The +// lookup column is interpolated into SQL, so an unlisted name must be refused before +// any query is built rather than trusted because call sites happen to pass literals. +func TestFindUserByFieldRecordRejectsUnlistedColumns(t *testing.T) { + db, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + if err := db.Table("w_users").Create(map[string]any{"id": 77, "username": "seeded"}).Error; err != nil { + t.Fatalf("seed user failed: %v", err) + } + + repository.SetDBServiceForTest(stubDBService{db: db}) + t.Cleanup(func() { repository.SetDBServiceForTest(nil) }) + + ctx := context.Background() + + user, err := repository.FindUserByFieldRecord(ctx, "username", "seeded") + if err != nil { + t.Fatalf("allowlisted lookup by username failed: %v", err) + } + if user.ID != 77 { + t.Errorf("allowlisted lookup returned ID %d, want 77", user.ID) + } + if _, err := repository.FindUserByFieldRecord(ctx, "id", uint64(77)); err != nil { + t.Errorf("allowlisted lookup by id failed: %v", err) + } + + cases := []struct { + name string + field string + }{ + {"tautology injection", `username = '' OR 1=1 --`}, + {"stacked statement", "id; DROP TABLE w_users"}, + {"column outside allow-list", "password"}, + {"empty field", ""}, + } + for _, tc := range cases { + if _, err := repository.FindUserByFieldRecord(ctx, tc.field, "seeded"); !errors.Is(err, errs.ErrUnsupportedUserLookupField) { + t.Errorf("%s: got err %v, want ErrUnsupportedUserLookupField", tc.name, err) + } + } + + var remaining int64 + if err := db.Table("w_users").Count(&remaining).Error; err != nil || remaining != 1 { + t.Fatalf("w_users damaged by rejected lookups: count=%d err=%v", remaining, err) + } +}