mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 14:06:36 +08:00
fix(user): clear need_change_password and invalidate cache on password change
This commit is contained in:
@@ -5,8 +5,10 @@ package cmd
|
||||
|
||||
import (
|
||||
"Wavelet/core"
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/alicebob/miniredis/v2"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
@@ -106,12 +108,20 @@ func TestNewWaveletAppProfiles(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestNewWaveletAppWithRedisEnabled(t *testing.T) {
|
||||
mr, err := miniredis.Run()
|
||||
require.NoError(t, err)
|
||||
defer mr.Close()
|
||||
|
||||
app := newWaveletApp(core.ProfileAll, core.WithConfigValues(map[string]any{
|
||||
"redis": map[string]any{
|
||||
"enabled": true,
|
||||
"addrs": []string{mr.Addr()},
|
||||
},
|
||||
}))
|
||||
require.NotNil(t, app)
|
||||
defer func() {
|
||||
_ = app.Stop(context.Background())
|
||||
}()
|
||||
require.NoError(t, app.Reconcile())
|
||||
|
||||
f, ok := app.Fiber("cache")
|
||||
|
||||
@@ -13,11 +13,30 @@ import (
|
||||
"encoding/hex"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"sync"
|
||||
|
||||
"github.com/gin-contrib/sessions"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
var (
|
||||
authMu sync.RWMutex
|
||||
authSvc contracts.AuthService
|
||||
)
|
||||
|
||||
// SetAuthService binds the AuthService for cache synchronization.
|
||||
func SetAuthService(s contracts.AuthService) {
|
||||
authMu.Lock()
|
||||
defer authMu.Unlock()
|
||||
authSvc = s
|
||||
}
|
||||
|
||||
func getAuthService() contracts.AuthService {
|
||||
authMu.RLock()
|
||||
defer authMu.RUnlock()
|
||||
return authSvc
|
||||
}
|
||||
|
||||
func getUserIDFromSession(c *gin.Context) uint64 {
|
||||
defer func() { _ = recover() }()
|
||||
session := sessions.Default(c)
|
||||
@@ -47,14 +66,15 @@ func getUserIDFromSession(c *gin.Context) uint64 {
|
||||
}
|
||||
|
||||
func invalidateUserCache(ctx context.Context, userID uint64) {
|
||||
// Cache invalidation delegated to AuthService via IoC at plugin Apply time.
|
||||
_ = ctx
|
||||
_ = userID
|
||||
if s := getAuthService(); s != nil {
|
||||
s.InvalidateCachedUser(ctx, userID)
|
||||
}
|
||||
}
|
||||
|
||||
func invalidateTokenCache(ctx context.Context, tokenHash string) {
|
||||
_ = ctx
|
||||
_ = tokenHash
|
||||
if s := getAuthService(); s != nil {
|
||||
s.InvalidateCachedToken(ctx, tokenHash)
|
||||
}
|
||||
}
|
||||
|
||||
// Login handles username and password authentication.
|
||||
@@ -171,6 +191,12 @@ func ChangePassword(c *gin.Context) {
|
||||
}
|
||||
invalidateUserCache(c.Request.Context(), user.ID)
|
||||
|
||||
sess := sessions.Default(c)
|
||||
sess.Set("need_change_password", false)
|
||||
if err := sess.Save(); err != nil {
|
||||
logger.ErrorF(c.Request.Context(), "save session failed on change-password: %v", err)
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, response.OKNil())
|
||||
}
|
||||
|
||||
|
||||
@@ -89,18 +89,27 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
||||
return nil
|
||||
})
|
||||
|
||||
// 0.1 Resolve auth service for middleware (via IoC, not direct import)
|
||||
// 0.1 Resolve auth service for middleware and cache invalidation (via IoC, not direct import)
|
||||
denyAuth := ginutil.AuthUnavailable()
|
||||
loginMW := denyAuth
|
||||
noTokenMW := denyAuth
|
||||
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
|
||||
SetAuthService(authSvc)
|
||||
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
|
||||
loginMW = mw
|
||||
}
|
||||
if mw, ok := authSvc.DisallowTokenAuthMiddleware().(gin.HandlerFunc); ok {
|
||||
noTokenMW = mw
|
||||
}
|
||||
} else {
|
||||
core.When[contracts.AuthService](ctx, func(svc contracts.AuthService) {
|
||||
SetAuthService(svc)
|
||||
})
|
||||
}
|
||||
ctx.OnDispose(func() error {
|
||||
SetAuthService(nil)
|
||||
return nil
|
||||
})
|
||||
|
||||
// 1. Register migrations
|
||||
ctx.Migrations().Register("user", userMigrations)
|
||||
|
||||
@@ -177,4 +177,24 @@ func TestUserLoginHTTPHandler(t *testing.T) {
|
||||
assert.Equal(t, http.StatusOK, wPlain.Code)
|
||||
assert.Contains(t, wPlain.Body.String(), `"username":"plain_admin"`)
|
||||
assert.Contains(t, wPlain.Body.String(), `"need_change_password":true`)
|
||||
cookieHeader := wPlain.Header().Get("Set-Cookie")
|
||||
assert.NotEmpty(t, cookieHeader)
|
||||
|
||||
// Change password
|
||||
r.POST("/api/v1/user/change-password", user.ChangePassword)
|
||||
changeBody := `{"old_password":"12345678","new_password":"NewStrongPassword123!"}`
|
||||
reqChange, _ := http.NewRequest(http.MethodPost, "/api/v1/user/change-password", bytes.NewBufferString(changeBody))
|
||||
reqChange.Header.Set("Content-Type", "application/json")
|
||||
reqChange.Header.Set("Cookie", cookieHeader)
|
||||
wChange := httptest.NewRecorder()
|
||||
|
||||
r.ServeHTTP(wChange, reqChange)
|
||||
assert.Equal(t, http.StatusOK, wChange.Code)
|
||||
|
||||
// Verify updated user model has encrypted password and need_change_password is false
|
||||
updatedUser, err := user.GetUserByUsername(context.Background(), "plain_admin")
|
||||
require.NoError(t, err)
|
||||
assert.False(t, updatedUser.IsPlaintextPassword())
|
||||
assert.True(t, updatedUser.CheckPassword("NewStrongPassword123!"))
|
||||
assert.False(t, updatedUser.NeedChangePassword)
|
||||
}
|
||||
|
||||
@@ -174,10 +174,14 @@ func (s *userServiceImpl) UpdatePassword(ctx context.Context, id uint64, oldPass
|
||||
return err
|
||||
}
|
||||
|
||||
return updateUserColumns(ctx, id, map[string]any{
|
||||
if err := updateUserColumns(ctx, id, map[string]any{
|
||||
"password": user.Password,
|
||||
columnUpdatedAt: time.Now(),
|
||||
})
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
invalidateUserCache(ctx, id)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *userServiceImpl) VerifyPassword(ctx context.Context, id uint64, password string) bool {
|
||||
|
||||
Reference in New Issue
Block a user