From 50f39a6983b4b20341383a93be82c89d6fc81605 Mon Sep 17 00:00:00 2001 From: ryan Date: Tue, 9 Jun 2026 20:07:25 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BF=AE=E5=A4=8D=E5=88=9B=E5=BB=BA=E8=B4=A6?= =?UTF-8?q?=E5=8F=B7=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/apps/user/routers.go | 19 +-- internal/apps/user/routers_test.go | 227 +++++++++++++++++++++++++++++ 2 files changed, 234 insertions(+), 12 deletions(-) create mode 100644 internal/apps/user/routers_test.go diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index 0415f443..ce188f64 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -30,6 +30,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/common" "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/task" "github.com/Rain-kl/Wavelet/internal/util" @@ -279,19 +280,12 @@ func Login(c *gin.Context) { } session := sessions.Default(c) - needChangePassword := false + needChangePassword := isPlaintext - // 如果是以明文密码登录,在数据库中置换为加密密码 if isPlaintext { - if err := user.SetEncryptedPassword(req.Password); err == nil { - if err := db.DB(ctx).Model(&user).Update("password", user.Password).Error; err != nil { - c.JSON(http.StatusOK, util.Err(errPasswordUpgradeFailed)) - return - } - needChangePassword = true - session.Set("need_change_password", true) - _ = session.Save() - } + session.Set("need_change_password", true) + } else { + session.Delete("need_change_password") } user.LastLoginAt = time.Now() @@ -358,6 +352,7 @@ func Register(c *gin.Context) { } user := model.User{ + ID: idgen.NextUint64ID(), Username: req.Username, Nickname: req.Nickname, Email: req.Email, @@ -372,7 +367,7 @@ func Register(c *gin.Context) { if user.Nickname == "" { user.Nickname = req.Username } - if err := user.SetPassword(req.Password); err != nil { + if err := user.SetEncryptedPassword(req.Password); err != nil { c.JSON(http.StatusOK, util.Err(err.Error())) return } diff --git a/internal/apps/user/routers_test.go b/internal/apps/user/routers_test.go new file mode 100644 index 00000000..bd96ef18 --- /dev/null +++ b/internal/apps/user/routers_test.go @@ -0,0 +1,227 @@ +/* +Copyright 2026 Arctel.net + +Licensed under the Apache License, Version 2.0 (the "License"); +you may not use this file except in compliance with the License. +You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software +distributed under the License is distributed on an "AS IS" BASIS, +WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +See the License for the specific language governing permissions and +limitations under the License. +*/ + +package user + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/apps/oauth" + "github.com/Rain-kl/Wavelet/internal/config" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/testhelper" + "github.com/Rain-kl/Wavelet/internal/util" + "github.com/gin-contrib/sessions" + "github.com/gin-contrib/sessions/cookie" + "github.com/gin-gonic/gin" +) + +func setupUserTestRouter(t *testing.T) *gin.Engine { + t.Helper() + + oldCookieName := config.Config.App.SessionCookieName + oldSecret := config.Config.App.SessionSecret + oldDomain := config.Config.App.SessionDomain + oldSecure := config.Config.App.SessionSecure + oldHTTPOnly := config.Config.App.SessionHTTPOnly + t.Cleanup(func() { + config.Config.App.SessionCookieName = oldCookieName + config.Config.App.SessionSecret = oldSecret + config.Config.App.SessionDomain = oldDomain + config.Config.App.SessionSecure = oldSecure + config.Config.App.SessionHTTPOnly = oldHTTPOnly + }) + + config.Config.App.SessionCookieName = "test_session_id" + config.Config.App.SessionSecret = "test_session_secret" + config.Config.App.SessionDomain = "" + config.Config.App.SessionSecure = false + config.Config.App.SessionHTTPOnly = true + + gin.SetMode(gin.TestMode) + r := gin.New() + store := cookie.NewStore([]byte(config.Config.App.SessionSecret)) + store.Options(util.GetSessionOptions(3600)) + r.Use(sessions.Sessions(config.Config.App.SessionCookieName, store)) + + api := r.Group("/api/v1") + api.POST("/user/register", Register) + api.POST("/user/login", Login) + api.GET("/user-info", oauth.LoginRequired(), oauth.UserInfo) + return r +} + +func performUserRequest(r http.Handler, method, path string, body []byte, cookies []*http.Cookie) *httptest.ResponseRecorder { + var reader *bytes.Reader + if body != nil { + reader = bytes.NewReader(body) + } else { + reader = bytes.NewReader(nil) + } + + req, _ := http.NewRequest(method, path, reader) + req.Header.Set("Content-Type", "application/json") + for _, c := range cookies { + req.AddCookie(c) + } + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + return w +} + +func sessionCookieFromResponse(t *testing.T, w *httptest.ResponseRecorder) *http.Cookie { + t.Helper() + + for _, c := range w.Result().Cookies() { + if c.Name == config.Config.App.SessionCookieName { + return c + } + } + t.Fatalf("sessionCookieFromResponse() did not find %q cookie", config.Config.App.SessionCookieName) + return nil +} + +func basicUserInfoFromResponse(t *testing.T, w *httptest.ResponseRecorder) oauth.BasicUserInfo { + t.Helper() + + var resp util.ResponseAny + if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { + t.Fatalf("basicUserInfoFromResponse() decode response failed: %v", err) + } + if resp.ErrorMsg != "" { + t.Fatalf("basicUserInfoFromResponse() error_msg = %q, want empty", resp.ErrorMsg) + } + data, _ := json.Marshal(resp.Data) + var info oauth.BasicUserInfo + if err := json.Unmarshal(data, &info); err != nil { + t.Fatalf("basicUserInfoFromResponse() decode data failed: %v", err) + } + return info +} + +func TestRegisterCreatesAuthenticatedEncryptedUser(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + router := setupUserTestRouter(t) + payload := registerRequest{ + Username: "newuser", + Password: "newpassword123", + Nickname: "New User", + } + body, _ := json.Marshal(payload) + + w := performUserRequest(router, http.MethodPost, "/api/v1/user/register", body, nil) + if w.Code != http.StatusOK { + t.Fatalf("Register(%q) status = %d, want %d. Body: %s", payload.Username, w.Code, http.StatusOK, w.Body.String()) + } + info := basicUserInfoFromResponse(t, w) + if info.NeedChangePassword { + t.Errorf("Register(%q) need_change_password = true, want false", payload.Username) + } + + var dbUser model.User + if err := dbConn.Where("username = ?", payload.Username).First(&dbUser).Error; err != nil { + t.Fatalf("Register(%q) db lookup failed: %v", payload.Username, err) + } + if dbUser.ID < 1000 { + t.Errorf("Register(%q) user ID = %d, want generated snowflake ID", payload.Username, dbUser.ID) + } + if !dbUser.IsPasswordEncrypted() { + t.Errorf("Register(%q) stored plaintext password, want encrypted password", payload.Username) + } + if !dbUser.CheckPassword(payload.Password) { + t.Errorf("Register(%q) stored password does not match original password", payload.Username) + } + + sessionCookie := sessionCookieFromResponse(t, w) + w = performUserRequest(router, http.MethodGet, "/api/v1/user-info", nil, []*http.Cookie{sessionCookie}) + if w.Code != http.StatusOK { + t.Fatalf("UserInfo() after Register(%q) status = %d, want %d. Body: %s", payload.Username, w.Code, http.StatusOK, w.Body.String()) + } +} + +func TestLoginRequiresPasswordChangeForInitialPlaintextAdmin(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + const ( + adminID = uint64(1) + adminUsername = "admin" + adminPassword = "12345678" + ) + now := time.Now() + if err := dbConn.Exec( + `INSERT INTO users (id, username, password, nickname, is_active, is_admin, last_login_at, created_at, updated_at) +VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, + adminID, + adminUsername, + adminPassword, + "Administrator", + true, + true, + now, + now, + now, + ).Error; err != nil { + t.Fatalf("seed initial admin failed: %v", err) + } + + router := setupUserTestRouter(t) + payload := loginRequest{ + Username: adminUsername, + Password: adminPassword, + } + body, _ := json.Marshal(payload) + + w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil) + if w.Code != http.StatusOK { + t.Fatalf("Login(%q) status = %d, want %d. Body: %s", adminUsername, w.Code, http.StatusOK, w.Body.String()) + } + info := basicUserInfoFromResponse(t, w) + if !info.NeedChangePassword { + t.Errorf("Login(%q) need_change_password = false, want true", adminUsername) + } + + var dbUser model.User + if err := dbConn.Where("username = ?", adminUsername).First(&dbUser).Error; err != nil { + t.Fatalf("Login(%q) db lookup failed: %v", adminUsername, err) + } + if dbUser.ID != adminID { + t.Errorf("Login(%q) user ID = %d, want %d", adminUsername, dbUser.ID, adminID) + } + if dbUser.IsPasswordEncrypted() { + t.Errorf("Login(%q) encrypted password during login, want plaintext until password change", adminUsername) + } + if !dbUser.CheckPassword(adminPassword) { + t.Errorf("Login(%q) stored password does not match original password", adminUsername) + } + + sessionCookie := sessionCookieFromResponse(t, w) + w = performUserRequest(router, http.MethodGet, "/api/v1/user-info", nil, []*http.Cookie{sessionCookie}) + if w.Code != http.StatusOK { + t.Fatalf("UserInfo() after Login(%q) status = %d, want %d. Body: %s", adminUsername, w.Code, http.StatusOK, w.Body.String()) + } + info = basicUserInfoFromResponse(t, w) + if !info.NeedChangePassword { + t.Errorf("UserInfo() after Login(%q) need_change_password = false, want true", adminUsername) + } +}