diff --git a/backend/cmd/app_test.go b/backend/cmd/app_test.go index 2625dea2..1ec92559 100644 --- a/backend/cmd/app_test.go +++ b/backend/cmd/app_test.go @@ -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") diff --git a/backend/plugins/domain/user/handlers.go b/backend/plugins/domain/user/handlers.go index 7755f740..2db068c2 100644 --- a/backend/plugins/domain/user/handlers.go +++ b/backend/plugins/domain/user/handlers.go @@ -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()) } diff --git a/backend/plugins/domain/user/plugin.go b/backend/plugins/domain/user/plugin.go index 7ba1bcd2..b0a17af5 100644 --- a/backend/plugins/domain/user/plugin.go +++ b/backend/plugins/domain/user/plugin.go @@ -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) diff --git a/backend/plugins/domain/user/plugin_test.go b/backend/plugins/domain/user/plugin_test.go index eac4ec15..a5a03783 100644 --- a/backend/plugins/domain/user/plugin_test.go +++ b/backend/plugins/domain/user/plugin_test.go @@ -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) } diff --git a/backend/plugins/domain/user/service.go b/backend/plugins/domain/user/service.go index 24753636..ee31ed04 100644 --- a/backend/plugins/domain/user/service.go +++ b/backend/plugins/domain/user/service.go @@ -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 {