diff --git a/AGENTS.md b/AGENTS.md index 8997eea6..c75d9438 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -66,6 +66,7 @@ Strong success criteria let you loop independently. Weak criteria ("make it work ## Git 提交规范 +每次完成一个功能点开发或修复一个问题后,务必提交 Git commit , 禁止推送远程仓库。 遵循 Conventional Commits:`(): `(例:`feat(auth): support email login`)。 ## 务必阅读匹配的 Skill diff --git a/backend/plugins/domain/auth/models.go b/backend/plugins/domain/auth/models.go index 74289f2f..00bb87aa 100644 --- a/backend/plugins/domain/auth/models.go +++ b/backend/plugins/domain/auth/models.go @@ -146,7 +146,7 @@ type CallbackRequest struct { // BasicUserInfo 用户基本信息结构体 type BasicUserInfo struct { - ID uint64 `json:"id"` + ID uint64 `json:"id,string"` Username string `json:"username"` Nickname string `json:"nickname"` Email string `json:"email"` diff --git a/backend/plugins/domain/auth/parse_userid_test.go b/backend/plugins/domain/auth/parse_userid_test.go new file mode 100644 index 00000000..212f10c0 --- /dev/null +++ b/backend/plugins/domain/auth/parse_userid_test.go @@ -0,0 +1,44 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "encoding/json" + "strconv" + "testing" +) + +func TestParseUserIDSnowflakeStringPreservesValue(t *testing.T) { + id := uint64(99835970421002240) + got := ParseUserID(strconv.FormatUint(id, 10)) + if got != id { + t.Errorf("ParseUserID(%q) = %d, want %d", strconv.FormatUint(id, 10), got, id) + } +} + +func TestParseUserIDJSONNumberAboveMaxSafeInteger(t *testing.T) { + // 2^53+1 cannot be represented as a distinct IEEE-754 float64. + id := uint64(9007199254740993) + got := ParseUserID(float64(id)) + if got == id { + t.Fatalf("ParseUserID(float64(%d)) = %d, want a rounded value", id, got) + } +} + +func TestBasicUserInfoJSONEncodesIDAsString(t *testing.T) { + info := BasicUserInfo{ID: 99835970421002240, Username: "plain_user"} + raw, err := json.Marshal(info) + if err != nil { + t.Fatalf("json.Marshal(BasicUserInfo) error = %v", err) + } + var probe struct { + ID json.RawMessage `json:"id"` + } + if err := json.Unmarshal(raw, &probe); err != nil { + t.Fatalf("json.Unmarshal probe error = %v", err) + } + if len(probe.ID) == 0 || probe.ID[0] != '"' { + t.Errorf("BasicUserInfo id JSON = %s, want a JSON string", probe.ID) + } +} diff --git a/backend/plugins/domain/auth/session.go b/backend/plugins/domain/auth/session.go index 97141413..007b375c 100644 --- a/backend/plugins/domain/auth/session.go +++ b/backend/plugins/domain/auth/session.go @@ -128,7 +128,7 @@ func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDT session.Clear() rotateSessionID(session) - session.Set(UserIDKey, user.ID) + session.Set(UserIDKey, strconv.FormatUint(user.ID, 10)) session.Set(UserNameKey, user.Username) if len(extras) > 0 { for key, value := range extras[0] { diff --git a/backend/plugins/domain/user/handlers.go b/backend/plugins/domain/user/handlers.go index 75b9ace0..e3ec20a6 100644 --- a/backend/plugins/domain/user/handlers.go +++ b/backend/plugins/domain/user/handlers.go @@ -108,7 +108,7 @@ func Login(c *gin.Context) { } sess := sessions.Default(c) - sess.Set(contracts.AuthUserIDKey, user.ID) + sess.Set(contracts.AuthUserIDKey, strconv.FormatUint(user.ID, 10)) sess.Set(contracts.AuthUserNameKey, user.Username) needChange := user.NeedChangePassword || user.IsPlaintextPassword() user.NeedChangePassword = needChange @@ -154,7 +154,7 @@ func Register(c *gin.Context) { } sess := sessions.Default(c) - sess.Set(contracts.AuthUserIDKey, newUser.ID) + sess.Set(contracts.AuthUserIDKey, strconv.FormatUint(newUser.ID, 10)) sess.Set(contracts.AuthUserNameKey, newUser.Username) if err := sess.Save(); err != nil { logger.ErrorF(c.Request.Context(), "save session failed on register: %v", err) diff --git a/backend/plugins/domain/user/repository.go b/backend/plugins/domain/user/repository.go index 5792c2ca..b77a4e0a 100644 --- a/backend/plugins/domain/user/repository.go +++ b/backend/plugins/domain/user/repository.go @@ -96,27 +96,6 @@ func UpdateUser(ctx context.Context, u *User) error { return getDB(ctx).Save(u).Error } -// ListUsers 分页查询用户 -func ListUsers(ctx context.Context, page, pageSize int, keyword string) ([]*User, int64, error) { - db := getDB(ctx).Model(&User{}) - if keyword != "" { - escaped := util.EscapeLike(keyword) - db = db.Where("username LIKE ? ESCAPE '\\' OR nickname LIKE ? ESCAPE '\\' OR email LIKE ? ESCAPE '\\'", "%"+escaped+"%", "%"+escaped+"%", "%"+escaped+"%") - } - - var total int64 - if err := db.Count(&total).Error; err != nil { - return nil, 0, err - } - - var users []*User - offset := (page - 1) * pageSize - if err := db.Offset(offset).Limit(pageSize).Order("id DESC").Find(&users).Error; err != nil { - return nil, 0, err - } - return users, total, nil -} - // GetAccessTokenByHash 通过 Hash 查询访问令牌 func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*AccessToken, error) { var token AccessToken diff --git a/backend/plugins/domain/user/session_access_test.go b/backend/plugins/domain/user/session_access_test.go new file mode 100644 index 00000000..09f92815 --- /dev/null +++ b/backend/plugins/domain/user/session_access_test.go @@ -0,0 +1,290 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user_test + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "Wavelet/pkg/response" + "Wavelet/plugins/domain/auth" + "Wavelet/plugins/domain/upload" + uploadmodels "Wavelet/plugins/domain/upload/models" + "Wavelet/plugins/domain/user" + "bytes" + "context" + "database/sql" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" + + database "Wavelet/plugins/infra/database" + "github.com/gin-contrib/sessions" + "github.com/gin-contrib/sessions/cookie" + "github.com/gin-gonic/gin" +) + +type stubStorageService struct{ contracts.StorageService } + +type loginEnvelope struct { + ErrorMsg string `json:"error_msg"` + Data json.RawMessage `json:"data"` +} + +func mountUserAuthEngine(t *testing.T) (*gin.Engine, contracts.UserService) { + t.Helper() + gin.SetMode(gin.TestMode) + + ctx := core.NewContext(context.Background()) + ctx.Config().SetSource(core.NewMapSource(nil)) + if err := ctx.Config().Resolve(); err != nil { + t.Fatalf("Config.Resolve() error = %v", err) + } + testDB := setupTestDB(t) + if err := testDB.AutoMigrate(&uploadmodels.Upload{}); err != nil { + t.Fatalf("AutoMigrate(Upload) error = %v", err) + } + if err := database.New(database.WithDB(testDB)).Apply(ctx); err != nil { + t.Fatalf("database.Apply() error = %v", err) + } + if err := auth.New().Apply(ctx); err != nil { + t.Fatalf("auth.Apply() error = %v", err) + } + if err := user.New().Apply(ctx); err != nil { + t.Fatalf("user.Apply() error = %v", err) + } + core.Provide[contracts.StorageService](ctx, stubStorageService{}) + if err := upload.New().Apply(ctx); err != nil { + t.Fatalf("upload.Apply() error = %v", err) + } + + userSvc, err := core.Inject[contracts.UserService](ctx) + if err != nil || userSvc == nil { + t.Fatalf("Inject UserService: svc=%v err=%v", userSvc, err) + } + + engine := gin.New() + engine.Use(response.ErrorHandlerMiddleware()) + engine.Use(sessions.Sessions("wavelet_session_id", cookie.NewStore([]byte("test-session-secret")))) + engine.Use(func(c *gin.Context) { + c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), ctx.Root())) + c.Next() + }) + + for _, rd := range ctx.Router().Routes() { + handlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers)) + for _, m := range rd.Middlewares { + h, ok := m.(gin.HandlerFunc) + if !ok { + fn, ok := m.(func(*gin.Context)) + if !ok { + t.Fatalf("unsupported middleware type %T for %s %s", m, rd.Method, rd.Path) + } + h = fn + } + handlers = append(handlers, h) + } + for _, raw := range rd.Handlers { + h, ok := raw.(gin.HandlerFunc) + if !ok { + fn, ok := raw.(func(*gin.Context)) + if !ok { + t.Fatalf("unsupported handler type %T for %s %s", raw, rd.Method, rd.Path) + } + h = fn + } + handlers = append(handlers, h) + } + engine.Handle(rd.Method, rd.Path, handlers...) + } + + return engine, userSvc +} + +func loginAndCookie(t *testing.T, engine *gin.Engine, username, password string) []*http.Cookie { + t.Helper() + body := `{"username":"` + username + `","password":"` + password + `"}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewBufferString(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("POST /api/v1/user/login username=%s status = %d, want 200 body=%s", username, rec.Code, rec.Body.String()) + } + cookies := rec.Result().Cookies() + if len(cookies) == 0 { + t.Fatalf("POST /api/v1/user/login username=%s Set-Cookie missing, headers=%v", username, rec.Header()) + } + return cookies +} + +func getWithCookies(engine *gin.Engine, path string, cookies []*http.Cookie) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodGet, path, nil) + for _, c := range cookies { + req.AddCookie(c) + } + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + return rec +} + +func TestNonAdminSessionCanAccessProtectedAPIs(t *testing.T) { + engine, userSvc := mountUserAuthEngine(t) + bg := context.Background() + + admin, err := userSvc.CreateUser(bg, contracts.CreateUserRequest{ + Username: "admin_user", + Password: "Password123!", + Email: "admin_user@example.com", + IsAdmin: true, + }) + if err != nil { + t.Fatalf("CreateUser(admin) error = %v", err) + } + if err := userSvc.SetUserAdmin(bg, admin.ID, true); err != nil { + t.Fatalf("SetUserAdmin() error = %v", err) + } + + member, err := userSvc.CreateUser(bg, contracts.CreateUserRequest{ + Username: "plain_user", + Password: "Password123!", + Email: "plain_user@example.com", + IsAdmin: false, + }) + if err != nil { + t.Fatalf("CreateUser(member) error = %v", err) + } + if member.IsAdmin { + t.Fatalf("CreateUser(member).IsAdmin = true, want false") + } + + cases := []struct { + name string + username string + wantID uint64 + }{ + {name: "admin", username: "admin_user", wantID: admin.ID}, + {name: "non-admin", username: "plain_user", wantID: member.ID}, + } + + protected := []string{ + "/api/v1/user/self", + "/api/v1/user-info", + "/api/v1/upload/my?page=1&page_size=12", + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cookies := loginAndCookie(t, engine, tc.username, "Password123!") + for _, path := range protected { + rec := getWithCookies(engine, path, cookies) + if rec.Code != http.StatusOK { + t.Errorf("GET %s as %s status = %d, want 200 body=%s", path, tc.name, rec.Code, rec.Body.String()) + continue + } + var env loginEnvelope + if err := json.Unmarshal(rec.Body.Bytes(), &env); err != nil { + t.Errorf("GET %s as %s decode error = %v body=%s", path, tc.name, err, rec.Body.String()) + continue + } + if env.ErrorMsg != "" { + t.Errorf("GET %s as %s error_msg = %q, want empty", path, tc.name, env.ErrorMsg) + } + if bytes.Contains(env.Data, []byte(`"username"`)) { + var payload struct { + ID json.RawMessage `json:"id"` + Username string `json:"username"` + } + if err := json.Unmarshal(env.Data, &payload); err != nil { + t.Errorf("GET %s as %s data decode error = %v data=%s", path, tc.name, err, string(env.Data)) + continue + } + if payload.Username != tc.username { + t.Errorf("GET %s as %s username = %q, want %q", path, tc.name, payload.Username, tc.username) + } + if len(payload.ID) == 0 || payload.ID[0] != '"' { + t.Errorf("GET %s as %s id JSON = %s, want a string (snowflake ids exceed JS MAX_SAFE_INTEGER)", path, tc.name, payload.ID) + } + } + } + }) + } +} + +func TestLoginBackfillsNullUserIDSoProtectedAPIsSucceed(t *testing.T) { + engine, _ := mountUserAuthEngine(t) + db := database.DB(context.Background()) + if db == nil { + t.Fatal("database.DB() = nil, want the test database") + } + + legacy := user.User{ + Username: "legacy_zero", + Email: "legacy_zero@example.com", + IsActive: true, + } + if err := legacy.SetEncryptedPassword("Password123!"); err != nil { + t.Fatalf("SetEncryptedPassword() error = %v", err) + } + sqlDB, err := db.DB() + if err != nil { + t.Fatalf("db.DB() error = %v", err) + } + if _, err := sqlDB.Exec( + "INSERT INTO w_users (id, username, password, email, is_active, is_admin) VALUES (0, ?, ?, ?, 1, 0)", + legacy.Username, legacy.Password, legacy.Email, + ); err != nil { + t.Fatalf("INSERT legacy user error = %v", err) + } + var stored sql.NullInt64 + if err := sqlDB.QueryRow("SELECT id FROM w_users WHERE username = ?", legacy.Username).Scan(&stored); err != nil { + t.Fatalf("SELECT id error = %v", err) + } + if stored.Valid && stored.Int64 != 0 { + t.Fatalf("legacy user id = %d, want 0 or NULL to reproduce the 401", stored.Int64) + } + + cookies := loginAndCookie(t, engine, legacy.Username, "Password123!") + protected := []string{ + "/api/v1/user/self", + "/api/v1/user-info", + "/api/v1/upload/my?page=1&page_size=12", + } + for _, path := range protected { + rec := getWithCookies(engine, path, cookies) + if rec.Code != http.StatusOK { + t.Errorf("GET %s as legacy_zero status = %d, want 200 body=%s", path, rec.Code, rec.Body.String()) + continue + } + var env loginEnvelope + if err := json.Unmarshal(rec.Body.Bytes(), &env); err != nil { + t.Errorf("GET %s as legacy_zero decode error = %v body=%s", path, err, rec.Body.String()) + continue + } + if env.ErrorMsg != "" { + t.Errorf("GET %s as legacy_zero error_msg = %q, want empty", path, env.ErrorMsg) + } + } + + var backfilled uint64 + if err := db.Raw("SELECT id FROM w_users WHERE username = ?", legacy.Username).Scan(&backfilled).Error; err != nil { + t.Fatalf("SELECT backfilled id error = %v", err) + } + if backfilled == 0 { + t.Errorf("legacy user id after login = 0, want a snowflake id") + } +} + +func TestLoginRequiredRejectsMissingSession(t *testing.T) { + engine, _ := mountUserAuthEngine(t) + rec := getWithCookies(engine, "/api/v1/user/self", nil) + if rec.Code != http.StatusUnauthorized { + t.Errorf("GET /api/v1/user/self without cookie status = %d, want 401 body=%s", rec.Code, rec.Body.String()) + } + body, _ := io.ReadAll(rec.Body) + if !bytes.Contains(body, []byte("未登录")) && !bytes.Contains(body, []byte("用户不存在")) { + t.Errorf("GET /api/v1/user/self without cookie body = %s, want 未登录", body) + } +}