mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-06 23:56:37 +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 (
|
import (
|
||||||
"Wavelet/core"
|
"Wavelet/core"
|
||||||
|
"context"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
"github.com/alicebob/miniredis/v2"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
@@ -106,12 +108,20 @@ func TestNewWaveletAppProfiles(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestNewWaveletAppWithRedisEnabled(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{
|
app := newWaveletApp(core.ProfileAll, core.WithConfigValues(map[string]any{
|
||||||
"redis": map[string]any{
|
"redis": map[string]any{
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
|
"addrs": []string{mr.Addr()},
|
||||||
},
|
},
|
||||||
}))
|
}))
|
||||||
require.NotNil(t, app)
|
require.NotNil(t, app)
|
||||||
|
defer func() {
|
||||||
|
_ = app.Stop(context.Background())
|
||||||
|
}()
|
||||||
require.NoError(t, app.Reconcile())
|
require.NoError(t, app.Reconcile())
|
||||||
|
|
||||||
f, ok := app.Fiber("cache")
|
f, ok := app.Fiber("cache")
|
||||||
|
|||||||
@@ -13,11 +13,30 @@ import (
|
|||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strconv"
|
"strconv"
|
||||||
|
"sync"
|
||||||
|
|
||||||
"github.com/gin-contrib/sessions"
|
"github.com/gin-contrib/sessions"
|
||||||
"github.com/gin-gonic/gin"
|
"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 {
|
func getUserIDFromSession(c *gin.Context) uint64 {
|
||||||
defer func() { _ = recover() }()
|
defer func() { _ = recover() }()
|
||||||
session := sessions.Default(c)
|
session := sessions.Default(c)
|
||||||
@@ -47,14 +66,15 @@ func getUserIDFromSession(c *gin.Context) uint64 {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func invalidateUserCache(ctx context.Context, userID uint64) {
|
func invalidateUserCache(ctx context.Context, userID uint64) {
|
||||||
// Cache invalidation delegated to AuthService via IoC at plugin Apply time.
|
if s := getAuthService(); s != nil {
|
||||||
_ = ctx
|
s.InvalidateCachedUser(ctx, userID)
|
||||||
_ = userID
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func invalidateTokenCache(ctx context.Context, tokenHash string) {
|
func invalidateTokenCache(ctx context.Context, tokenHash string) {
|
||||||
_ = ctx
|
if s := getAuthService(); s != nil {
|
||||||
_ = tokenHash
|
s.InvalidateCachedToken(ctx, tokenHash)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Login handles username and password authentication.
|
// Login handles username and password authentication.
|
||||||
@@ -171,6 +191,12 @@ func ChangePassword(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
invalidateUserCache(c.Request.Context(), user.ID)
|
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())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -89,18 +89,27 @@ func (p *Plugin) Apply(ctx *core.Context) error {
|
|||||||
return nil
|
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()
|
denyAuth := ginutil.AuthUnavailable()
|
||||||
loginMW := denyAuth
|
loginMW := denyAuth
|
||||||
noTokenMW := denyAuth
|
noTokenMW := denyAuth
|
||||||
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
|
if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil {
|
||||||
|
SetAuthService(authSvc)
|
||||||
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
|
if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok {
|
||||||
loginMW = mw
|
loginMW = mw
|
||||||
}
|
}
|
||||||
if mw, ok := authSvc.DisallowTokenAuthMiddleware().(gin.HandlerFunc); ok {
|
if mw, ok := authSvc.DisallowTokenAuthMiddleware().(gin.HandlerFunc); ok {
|
||||||
noTokenMW = mw
|
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
|
// 1. Register migrations
|
||||||
ctx.Migrations().Register("user", userMigrations)
|
ctx.Migrations().Register("user", userMigrations)
|
||||||
|
|||||||
@@ -177,4 +177,24 @@ func TestUserLoginHTTPHandler(t *testing.T) {
|
|||||||
assert.Equal(t, http.StatusOK, wPlain.Code)
|
assert.Equal(t, http.StatusOK, wPlain.Code)
|
||||||
assert.Contains(t, wPlain.Body.String(), `"username":"plain_admin"`)
|
assert.Contains(t, wPlain.Body.String(), `"username":"plain_admin"`)
|
||||||
assert.Contains(t, wPlain.Body.String(), `"need_change_password":true`)
|
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 err
|
||||||
}
|
}
|
||||||
|
|
||||||
return updateUserColumns(ctx, id, map[string]any{
|
if err := updateUserColumns(ctx, id, map[string]any{
|
||||||
"password": user.Password,
|
"password": user.Password,
|
||||||
columnUpdatedAt: time.Now(),
|
columnUpdatedAt: time.Now(),
|
||||||
})
|
}); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
invalidateUserCache(ctx, id)
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *userServiceImpl) VerifyPassword(ctx context.Context, id uint64, password string) bool {
|
func (s *userServiceImpl) VerifyPassword(ctx context.Context, id uint64, password string) bool {
|
||||||
|
|||||||
Reference in New Issue
Block a user