Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 23 additions & 0 deletions models/services.go
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
package models

import "context"

type ServiceID string

const (
Expand Down Expand Up @@ -29,3 +31,24 @@ type ServiceRegistry interface {
Register(name string, service any)
Get(name string) any
}

const ContextServiceRegistry ContextKey = "service.service_registry"

func NewContextWithServiceRegistry(ctx context.Context, registry ServiceRegistry) context.Context {
return context.WithValue(ctx, ContextServiceRegistry, registry)
}

func GetServiceRegistry(ctx context.Context) ServiceRegistry {
registry, _ := ctx.Value(ContextServiceRegistry).(ServiceRegistry)
return registry
}

func GetServiceFromContext[T any](ctx context.Context, id ServiceID) (bool, T) {
registry := GetServiceRegistry(ctx)
if registry == nil {
var zero T
return false, zero
}
svc, ok := registry.Get(id.String()).(T)
return ok, svc
}
2 changes: 1 addition & 1 deletion plugins/email-password/plugin.go
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,7 @@ func (p *EmailPasswordPlugin) Init(ctx *models.PluginContext) error {
}
p.emailTemplateManager = emailTemplateManager

p.hooksExecutor = services.NewServiceHookExecutor(p.pluginConfig.ServiceHooks, p.logger)
p.hooksExecutor = services.NewServiceHookExecutor(p.pluginConfig.ServiceHooks, p.logger, ctx.ServiceRegistry)

p.Api = BuildAPI(p)

Expand Down
19 changes: 15 additions & 4 deletions plugins/email-password/services/hooks_executor.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,25 +8,28 @@ import (
)

type ServiceHookExecutor struct {
config *types.EmailPasswordServiceHooksConfig
logger models.Logger
config *types.EmailPasswordServiceHooksConfig
logger models.Logger
registry models.ServiceRegistry
}

func NewServiceHookExecutor(config *types.EmailPasswordServiceHooksConfig, logger models.Logger) *ServiceHookExecutor {
return &ServiceHookExecutor{config: config, logger: logger}
func NewServiceHookExecutor(config *types.EmailPasswordServiceHooksConfig, logger models.Logger, registry models.ServiceRegistry) *ServiceHookExecutor {
return &ServiceHookExecutor{config: config, logger: logger, registry: registry}
}

func (e *ServiceHookExecutor) BeforeSignUp(ctx context.Context, user *models.User) error {
if e == nil || e.config == nil || e.config.SignUp == nil || e.config.SignUp.BeforeSignUp == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
return e.config.SignUp.BeforeSignUp(ctx, user)
}

func (e *ServiceHookExecutor) AfterSignUp(ctx context.Context, result *types.SignUpResult) error {
if e == nil || e.config == nil || e.config.SignUp == nil || e.config.SignUp.AfterSignUp == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
if err := e.config.SignUp.AfterSignUp(ctx, result); err != nil {
e.logger.Error("after sign up hook failed", "error", err.Error())
}
Expand All @@ -37,13 +40,15 @@ func (e *ServiceHookExecutor) BeforeSignIn(ctx context.Context, user *models.Use
if e == nil || e.config == nil || e.config.SignIn == nil || e.config.SignIn.BeforeSignIn == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
return e.config.SignIn.BeforeSignIn(ctx, user)
}

func (e *ServiceHookExecutor) AfterSignIn(ctx context.Context, result *types.SignInResult) error {
if e == nil || e.config == nil || e.config.SignIn == nil || e.config.SignIn.AfterSignIn == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
if err := e.config.SignIn.AfterSignIn(ctx, result); err != nil {
e.logger.Error("after sign in hook failed", "error", err.Error())
}
Expand All @@ -54,6 +59,7 @@ func (e *ServiceHookExecutor) AfterVerifyEmail(ctx context.Context, user *models
if e == nil || e.config == nil || e.config.EmailVerification == nil || e.config.EmailVerification.AfterVerifyEmail == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
if err := e.config.EmailVerification.AfterVerifyEmail(ctx, user, verificationType); err != nil {
e.logger.Error("after verify email hook failed", "error", err.Error())
}
Expand All @@ -64,20 +70,23 @@ func (e *ServiceHookExecutor) BeforeRequestPasswordReset(ctx context.Context, us
if e == nil || e.config == nil || e.config.PasswordReset == nil || e.config.PasswordReset.BeforeRequestPasswordReset == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
return e.config.PasswordReset.BeforeRequestPasswordReset(ctx, user)
}

func (e *ServiceHookExecutor) BeforeChangePassword(ctx context.Context, user *models.User, newPassword string) error {
if e == nil || e.config == nil || e.config.PasswordChange == nil || e.config.PasswordChange.BeforeChangePassword == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
return e.config.PasswordChange.BeforeChangePassword(ctx, user, newPassword)
}

func (e *ServiceHookExecutor) AfterChangePassword(ctx context.Context, user *models.User) error {
if e == nil || e.config == nil || e.config.PasswordChange == nil || e.config.PasswordChange.AfterChangePassword == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
if err := e.config.PasswordChange.AfterChangePassword(ctx, user); err != nil {
e.logger.Error("after change password hook failed", "error", err.Error())
}
Expand All @@ -88,13 +97,15 @@ func (e *ServiceHookExecutor) BeforeRequestEmailChange(ctx context.Context, user
if e == nil || e.config == nil || e.config.EmailChange == nil || e.config.EmailChange.BeforeRequestEmailChange == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
return e.config.EmailChange.BeforeRequestEmailChange(ctx, user)
}

func (e *ServiceHookExecutor) AfterEmailChanged(ctx context.Context, user *models.User, oldEmail, newEmail string) error {
if e == nil || e.config == nil || e.config.EmailChange == nil || e.config.EmailChange.AfterEmailChanged == nil {
return nil
}
ctx = models.NewContextWithServiceRegistry(ctx, e.registry)
if err := e.config.EmailChange.AfterEmailChanged(ctx, user, oldEmail, newEmail); err != nil {
e.logger.Error("after email changed hook failed", "error", err.Error())
}
Expand Down
16 changes: 8 additions & 8 deletions plugins/email-password/services/hooks_executor_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@ import (
func TestServiceHookExecutor_NilConfigIsNoop(t *testing.T) {
t.Parallel()

executor := NewServiceHookExecutor(nil, nil)
executor := NewServiceHookExecutor(nil, nil, nil)
ctx := context.Background()
user := &models.User{ID: "user-1", Email: "test@example.com"}
signUpResult := &types.SignUpResult{User: user}
Expand Down Expand Up @@ -99,7 +99,7 @@ func TestServiceHookExecutor_SignUpHooks(t *testing.T) {
return nil
},
},
}, nil)
}, nil, nil)

ctx := context.Background()
user := &models.User{ID: "user-1", Email: "test@example.com"}
Expand Down Expand Up @@ -143,7 +143,7 @@ func TestServiceHookExecutor_SignInHooks(t *testing.T) {
return nil
},
},
}, nil)
}, nil, nil)

ctx := context.Background()
user := &models.User{ID: "user-1", Email: "test@example.com"}
Expand Down Expand Up @@ -178,7 +178,7 @@ func TestServiceHookExecutor_AfterVerifyEmailHook(t *testing.T) {
return nil
},
},
}, nil)
}, nil, nil)

ctx := context.Background()
user := &models.User{ID: "user-1"}
Expand Down Expand Up @@ -209,7 +209,7 @@ func TestServiceHookExecutor_BeforeChangePasswordHook(t *testing.T) {
return nil
},
},
}, nil)
}, nil, nil)

ctx := context.Background()
user := &models.User{ID: "user-1"}
Expand Down Expand Up @@ -242,7 +242,7 @@ func TestServiceHookExecutor_AfterEmailChangedHook(t *testing.T) {
return nil
},
},
}, nil)
}, nil, nil)

ctx := context.Background()
user := &models.User{ID: "user-1"}
Expand Down Expand Up @@ -272,7 +272,7 @@ func TestServiceHookExecutor_BeforeHookError(t *testing.T) {
return someErr
},
},
}, nil)
}, nil, nil)

err := executor.BeforeSignUp(context.Background(), &models.User{ID: "user-1"})
if !errors.Is(err, someErr) {
Expand All @@ -290,7 +290,7 @@ func TestServiceHookExecutor_AfterHookErrorIsLoggedNotReturned(t *testing.T) {
return errors.New("hook error")
},
},
}, logger)
}, logger, nil)

err := executor.AfterSignUp(context.Background(), &types.SignUpResult{User: &models.User{ID: "user-1"}})
if err != nil {
Expand Down
6 changes: 3 additions & 3 deletions plugins/email-password/usecases/change_password_hooks_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ func TestChangePasswordUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.tokenSvc.On("Hash", mock.Anything).Return("hashed-token", nil)
f.verificationSvc.On("GetByToken", mock.Anything, "hashed-token").Return(&models.Verification{
ID: "ver-1", UserID: new("user-1"), Type: models.TypePasswordResetRequest, ExpiresAt: time.Now().Add(time.Hour), Identifier: "test@example.com",
Expand All @@ -59,7 +59,7 @@ func TestChangePasswordUseCaseHooks(t *testing.T) {
return errors.New("password in history")
},
},
}, f.logger)
}, f.logger, nil)
f.tokenSvc.On("Hash", mock.Anything).Return("hashed-token", nil)
f.verificationSvc.On("GetByToken", mock.Anything, "hashed-token").Return(&models.Verification{
ID: "ver-1", UserID: new("user-1"), Type: models.TypePasswordResetRequest, ExpiresAt: time.Now().Add(time.Hour), Identifier: "test@example.com",
Expand All @@ -82,7 +82,7 @@ func TestChangePasswordUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.tokenSvc.On("Hash", mock.Anything).Return("hashed-token", nil)
f.verificationSvc.On("GetByToken", mock.Anything, "hashed-token").Return(&models.Verification{
ID: "ver-1", UserID: new("user-1"), Type: models.TypePasswordResetRequest, ExpiresAt: time.Now().Add(time.Hour), Identifier: "test@example.com",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@ func TestRequestEmailChangeUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByID", mock.Anything, "user-1").Return(&models.User{ID: "user-1", Email: "old@example.com"}, nil)
f.userSvc.On("GetByEmail", mock.Anything, "new@example.com").Return(nil, nil)
f.tokenSvc.On("Generate").Return("token-123", nil)
Expand All @@ -52,7 +52,7 @@ func TestRequestEmailChangeUseCaseHooks(t *testing.T) {
return errors.New("email change not allowed")
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByID", mock.Anything, "user-1").Return(&models.User{ID: "user-1", Email: "old@example.com"}, nil)
f.userSvc.On("GetByEmail", mock.Anything, "new@example.com").Return(nil, nil)
},
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ func TestRequestPasswordResetUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "test@example.com").Return(&models.User{ID: "user-1", Email: "test@example.com"}, nil)
f.tokenSvc.On("Generate").Return("token-123", nil)
f.tokenSvc.On("Hash", mock.Anything).Return("hashed-token", nil)
Expand All @@ -53,7 +53,7 @@ func TestRequestPasswordResetUseCaseHooks(t *testing.T) {
return errors.New("rate limited")
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "test@example.com").Return(&models.User{ID: "user-1", Email: "test@example.com"}, nil)
},
assert: func(t *testing.T, err error) {
Expand All @@ -72,7 +72,7 @@ func TestRequestPasswordResetUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "nonexistent@example.com").Return(nil, nil)
},
assert: func(t *testing.T, err error) {
Expand Down
8 changes: 4 additions & 4 deletions plugins/email-password/usecases/sign_in_hooks_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ func TestSignInUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "test@example.com").Return(&models.User{ID: "user-1", Email: "test@example.com"}, nil)
f.accountSvc.On("GetByUserIDAndProvider", mock.Anything, "user-1", mock.Anything).Return(&models.Account{ID: "acc-1", Password: new("hashed")}, nil)
f.passwordSvc.On("Verify", mock.Anything, mock.Anything).Return(true)
Expand All @@ -57,7 +57,7 @@ func TestSignInUseCaseHooks(t *testing.T) {
return errors.New("hook rejected")
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "test@example.com").Return(&models.User{ID: "user-1", Email: "test@example.com"}, nil)
},
assert: func(t *testing.T, result *types.SignInResult, err error) {
Expand All @@ -77,7 +77,7 @@ func TestSignInUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "nonexistent@example.com").Return(nil, nil)
},
assert: func(t *testing.T, result *types.SignInResult, err error) {
Expand All @@ -97,7 +97,7 @@ func TestSignInUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "test@example.com").Return(&models.User{ID: "user-1", Email: "test@example.com"}, nil)
f.accountSvc.On("GetByUserIDAndProvider", mock.Anything, "user-1", mock.Anything).Return(&models.Account{ID: "acc-1", Password: new("hashed")}, nil)
f.passwordSvc.On("Verify", mock.Anything, mock.Anything).Return(true)
Expand Down
6 changes: 3 additions & 3 deletions plugins/email-password/usecases/sign_up_hooks_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ func TestSignUpUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "test@example.com").Return(nil, nil)
f.userSvc.On("Create", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(&models.User{ID: "user-1", Name: "Test User", Email: "test@example.com"}, nil)
f.passwordSvc.On("Hash", mock.Anything).Return("hashed", nil)
Expand All @@ -57,7 +57,7 @@ func TestSignUpUseCaseHooks(t *testing.T) {
return errors.New("hook rejected")
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "test@example.com").Return(nil, nil)
},
assert: func(t *testing.T, result *types.SignUpResult, err error) {
Expand All @@ -77,7 +77,7 @@ func TestSignUpUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.userSvc.On("GetByEmail", mock.Anything, "test@example.com").Return(nil, nil)
f.userSvc.On("Create", mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything, mock.Anything).Return(&models.User{ID: "user-1", Name: "Test User", Email: "test@example.com"}, nil)
f.passwordSvc.On("Hash", mock.Anything).Return("hashed", nil)
Expand Down
6 changes: 3 additions & 3 deletions plugins/email-password/usecases/verify_email_hooks_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ func TestVerifyEmailUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.tokenSvc.On("Hash", mock.Anything).Return("hashed-token", nil)
f.verificationSvc.On("GetByToken", mock.Anything, "hashed-token").Return(&models.Verification{
ID: "ver-1", UserID: new("user-1"), Type: models.TypeEmailVerification, ExpiresAt: time.Now().Add(time.Hour), Identifier: "test@example.com",
Expand All @@ -57,7 +57,7 @@ func TestVerifyEmailUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.tokenSvc.On("Hash", mock.Anything).Return("hashed-token", nil)
f.verificationSvc.On("GetByToken", mock.Anything, "hashed-token").Return(&models.Verification{
ID: "ver-1", UserID: new("user-1"), Type: models.TypePasswordResetRequest, ExpiresAt: time.Now().Add(time.Hour), Identifier: "test@example.com",
Expand All @@ -79,7 +79,7 @@ func TestVerifyEmailUseCaseHooks(t *testing.T) {
return nil
},
},
}, f.logger)
}, f.logger, nil)
f.tokenSvc.On("Hash", mock.Anything).Return("hashed-token", nil)
f.verificationSvc.On("GetByToken", mock.Anything, "hashed-token").Return(&models.Verification{
ID: "ver-1", UserID: new("user-1"), Type: models.TypeEmailResetRequest, ExpiresAt: time.Now().Add(time.Hour), Identifier: "new@example.com",
Expand Down
2 changes: 1 addition & 1 deletion plugins/organizations/plugin.go
Original file line number Diff line number Diff line change
Expand Up @@ -84,7 +84,7 @@ func (p *OrganizationsPlugin) Init(ctx *models.PluginContext) error {
return err
}

p.hooksExecutor = services.NewServiceHookExecutor(p.pluginConfig.ServiceHooks)
p.hooksExecutor = services.NewServiceHookExecutor(p.pluginConfig.ServiceHooks, ctx.ServiceRegistry)
p.organizationRepo = repositories.NewBunOrganizationRepository(ctx.DB)
p.invitationRepo = repositories.NewBunOrganizationInvitationRepository(ctx.DB)
p.memberRepo = repositories.NewBunOrganizationMemberRepository(ctx.DB)
Expand Down
Loading