diff --git a/models/services.go b/models/services.go index a3e642f..e9a837c 100644 --- a/models/services.go +++ b/models/services.go @@ -1,5 +1,7 @@ package models +import "context" + type ServiceID string const ( @@ -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 +} diff --git a/plugins/email-password/plugin.go b/plugins/email-password/plugin.go index dd2b722..5ff54e3 100644 --- a/plugins/email-password/plugin.go +++ b/plugins/email-password/plugin.go @@ -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) diff --git a/plugins/email-password/services/hooks_executor.go b/plugins/email-password/services/hooks_executor.go index e4ffeb5..4a7def2 100644 --- a/plugins/email-password/services/hooks_executor.go +++ b/plugins/email-password/services/hooks_executor.go @@ -8,18 +8,20 @@ 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) } @@ -27,6 +29,7 @@ func (e *ServiceHookExecutor) AfterSignUp(ctx context.Context, result *types.Sig 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()) } @@ -37,6 +40,7 @@ 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) } @@ -44,6 +48,7 @@ func (e *ServiceHookExecutor) AfterSignIn(ctx context.Context, result *types.Sig 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()) } @@ -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()) } @@ -64,6 +70,7 @@ 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) } @@ -71,6 +78,7 @@ func (e *ServiceHookExecutor) BeforeChangePassword(ctx context.Context, user *mo 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) } @@ -78,6 +86,7 @@ func (e *ServiceHookExecutor) AfterChangePassword(ctx context.Context, user *mod 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()) } @@ -88,6 +97,7 @@ 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) } @@ -95,6 +105,7 @@ func (e *ServiceHookExecutor) AfterEmailChanged(ctx context.Context, user *model 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()) } diff --git a/plugins/email-password/services/hooks_executor_test.go b/plugins/email-password/services/hooks_executor_test.go index dbd5e79..c7c0bf0 100644 --- a/plugins/email-password/services/hooks_executor_test.go +++ b/plugins/email-password/services/hooks_executor_test.go @@ -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} @@ -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"} @@ -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"} @@ -178,7 +178,7 @@ func TestServiceHookExecutor_AfterVerifyEmailHook(t *testing.T) { return nil }, }, - }, nil) + }, nil, nil) ctx := context.Background() user := &models.User{ID: "user-1"} @@ -209,7 +209,7 @@ func TestServiceHookExecutor_BeforeChangePasswordHook(t *testing.T) { return nil }, }, - }, nil) + }, nil, nil) ctx := context.Background() user := &models.User{ID: "user-1"} @@ -242,7 +242,7 @@ func TestServiceHookExecutor_AfterEmailChangedHook(t *testing.T) { return nil }, }, - }, nil) + }, nil, nil) ctx := context.Background() user := &models.User{ID: "user-1"} @@ -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) { @@ -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 { diff --git a/plugins/email-password/usecases/change_password_hooks_test.go b/plugins/email-password/usecases/change_password_hooks_test.go index 09a6a19..b7102ba 100644 --- a/plugins/email-password/usecases/change_password_hooks_test.go +++ b/plugins/email-password/usecases/change_password_hooks_test.go @@ -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", @@ -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", @@ -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", diff --git a/plugins/email-password/usecases/request_email_change_hooks_test.go b/plugins/email-password/usecases/request_email_change_hooks_test.go index 5c4c161..a11c6eb 100644 --- a/plugins/email-password/usecases/request_email_change_hooks_test.go +++ b/plugins/email-password/usecases/request_email_change_hooks_test.go @@ -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) @@ -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) }, diff --git a/plugins/email-password/usecases/request_password_reset_hooks_test.go b/plugins/email-password/usecases/request_password_reset_hooks_test.go index 7cb0c09..6355063 100644 --- a/plugins/email-password/usecases/request_password_reset_hooks_test.go +++ b/plugins/email-password/usecases/request_password_reset_hooks_test.go @@ -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) @@ -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) { @@ -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) { diff --git a/plugins/email-password/usecases/sign_in_hooks_test.go b/plugins/email-password/usecases/sign_in_hooks_test.go index 4e97db8..f4f16f2 100644 --- a/plugins/email-password/usecases/sign_in_hooks_test.go +++ b/plugins/email-password/usecases/sign_in_hooks_test.go @@ -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) @@ -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) { @@ -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) { @@ -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) diff --git a/plugins/email-password/usecases/sign_up_hooks_test.go b/plugins/email-password/usecases/sign_up_hooks_test.go index f728c3d..ad6de65 100644 --- a/plugins/email-password/usecases/sign_up_hooks_test.go +++ b/plugins/email-password/usecases/sign_up_hooks_test.go @@ -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) @@ -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) { @@ -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) diff --git a/plugins/email-password/usecases/verify_email_hooks_test.go b/plugins/email-password/usecases/verify_email_hooks_test.go index 9a0f8ad..39e723f 100644 --- a/plugins/email-password/usecases/verify_email_hooks_test.go +++ b/plugins/email-password/usecases/verify_email_hooks_test.go @@ -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", @@ -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", @@ -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", diff --git a/plugins/organizations/plugin.go b/plugins/organizations/plugin.go index c5b9359..453ec3b 100644 --- a/plugins/organizations/plugin.go +++ b/plugins/organizations/plugin.go @@ -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) diff --git a/plugins/organizations/services/hooks_executor.go b/plugins/organizations/services/hooks_executor.go index fdba76d..33d9b15 100644 --- a/plugins/organizations/services/hooks_executor.go +++ b/plugins/organizations/services/hooks_executor.go @@ -8,17 +8,19 @@ import ( ) type ServiceHookExecutor struct { - config *types.OrganizationsServiceHooksConfig + config *types.OrganizationsServiceHooksConfig + registry models.ServiceRegistry } -func NewServiceHookExecutor(config *types.OrganizationsServiceHooksConfig) *ServiceHookExecutor { - return &ServiceHookExecutor{config: config} +func NewServiceHookExecutor(config *types.OrganizationsServiceHooksConfig, registry models.ServiceRegistry) *ServiceHookExecutor { + return &ServiceHookExecutor{config: config, registry: registry} } func (e *ServiceHookExecutor) BeforeCreateOrganization(ctx context.Context, actor *models.Actor, organization *types.Organization) error { if e == nil || e.config == nil || e.config.Organizations == nil || e.config.Organizations.BeforeCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Organizations.BeforeCreate(ctx, actor, organization) } @@ -26,6 +28,7 @@ func (e *ServiceHookExecutor) AfterCreateOrganization(ctx context.Context, actor if e == nil || e.config == nil || e.config.Organizations == nil || e.config.Organizations.AfterCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Organizations.AfterCreate(ctx, actor, organization) } @@ -33,6 +36,7 @@ func (e *ServiceHookExecutor) BeforeUpdateOrganization(ctx context.Context, acto if e == nil || e.config == nil || e.config.Organizations == nil || e.config.Organizations.BeforeUpdate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Organizations.BeforeUpdate(ctx, actor, organization) } @@ -40,6 +44,7 @@ func (e *ServiceHookExecutor) AfterUpdateOrganization(ctx context.Context, actor if e == nil || e.config == nil || e.config.Organizations == nil || e.config.Organizations.AfterUpdate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Organizations.AfterUpdate(ctx, actor, organization) } @@ -47,6 +52,7 @@ func (e *ServiceHookExecutor) BeforeDeleteOrganization(ctx context.Context, acto if e == nil || e.config == nil || e.config.Organizations == nil || e.config.Organizations.BeforeDelete == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Organizations.BeforeDelete(ctx, actor, organization) } @@ -54,6 +60,7 @@ func (e *ServiceHookExecutor) AfterDeleteOrganization(ctx context.Context, actor if e == nil || e.config == nil || e.config.Organizations == nil || e.config.Organizations.AfterDelete == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Organizations.AfterDelete(ctx, actor, organization) } @@ -61,6 +68,7 @@ func (e *ServiceHookExecutor) BeforeCreateOrganizationMember(ctx context.Context if e == nil || e.config == nil || e.config.Members == nil || e.config.Members.BeforeCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Members.BeforeCreate(ctx, actor, member) } @@ -68,6 +76,7 @@ func (e *ServiceHookExecutor) AfterCreateOrganizationMember(ctx context.Context, if e == nil || e.config == nil || e.config.Members == nil || e.config.Members.AfterCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Members.AfterCreate(ctx, actor, member) } @@ -75,6 +84,7 @@ func (e *ServiceHookExecutor) BeforeUpdateOrganizationMember(ctx context.Context if e == nil || e.config == nil || e.config.Members == nil || e.config.Members.BeforeUpdate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Members.BeforeUpdate(ctx, actor, member) } @@ -82,6 +92,7 @@ func (e *ServiceHookExecutor) AfterUpdateOrganizationMember(ctx context.Context, if e == nil || e.config == nil || e.config.Members == nil || e.config.Members.AfterUpdate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Members.AfterUpdate(ctx, actor, member) } @@ -89,6 +100,7 @@ func (e *ServiceHookExecutor) BeforeDeleteOrganizationMember(ctx context.Context if e == nil || e.config == nil || e.config.Members == nil || e.config.Members.BeforeDelete == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Members.BeforeDelete(ctx, actor, member) } @@ -96,6 +108,7 @@ func (e *ServiceHookExecutor) AfterDeleteOrganizationMember(ctx context.Context, if e == nil || e.config == nil || e.config.Members == nil || e.config.Members.AfterDelete == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Members.AfterDelete(ctx, actor, member) } @@ -103,6 +116,7 @@ func (e *ServiceHookExecutor) BeforeCreateOrganizationInvitation(ctx context.Con if e == nil || e.config == nil || e.config.Invitations == nil || e.config.Invitations.BeforeCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Invitations.BeforeCreate(ctx, actor, invitation) } @@ -110,6 +124,7 @@ func (e *ServiceHookExecutor) AfterCreateOrganizationInvitation(ctx context.Cont if e == nil || e.config == nil || e.config.Invitations == nil || e.config.Invitations.AfterCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Invitations.AfterCreate(ctx, actor, invitation) } @@ -117,6 +132,7 @@ func (e *ServiceHookExecutor) BeforeUpdateOrganizationInvitation(ctx context.Con if e == nil || e.config == nil || e.config.Invitations == nil || e.config.Invitations.BeforeUpdate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Invitations.BeforeUpdate(ctx, actor, invitation) } @@ -124,6 +140,7 @@ func (e *ServiceHookExecutor) AfterUpdateOrganizationInvitation(ctx context.Cont if e == nil || e.config == nil || e.config.Invitations == nil || e.config.Invitations.AfterUpdate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Invitations.AfterUpdate(ctx, actor, invitation) } @@ -131,6 +148,7 @@ func (e *ServiceHookExecutor) BeforeCreateOrganizationTeam(ctx context.Context, if e == nil || e.config == nil || e.config.Teams == nil || e.config.Teams.BeforeCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Teams.BeforeCreate(ctx, actor, team) } @@ -138,6 +156,7 @@ func (e *ServiceHookExecutor) AfterCreateOrganizationTeam(ctx context.Context, a if e == nil || e.config == nil || e.config.Teams == nil || e.config.Teams.AfterCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Teams.AfterCreate(ctx, actor, team) } @@ -145,6 +164,7 @@ func (e *ServiceHookExecutor) BeforeUpdateOrganizationTeam(ctx context.Context, if e == nil || e.config == nil || e.config.Teams == nil || e.config.Teams.BeforeUpdate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Teams.BeforeUpdate(ctx, actor, team) } @@ -152,6 +172,7 @@ func (e *ServiceHookExecutor) AfterUpdateOrganizationTeam(ctx context.Context, a if e == nil || e.config == nil || e.config.Teams == nil || e.config.Teams.AfterUpdate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Teams.AfterUpdate(ctx, actor, team) } @@ -159,6 +180,7 @@ func (e *ServiceHookExecutor) BeforeDeleteOrganizationTeam(ctx context.Context, if e == nil || e.config == nil || e.config.Teams == nil || e.config.Teams.BeforeDelete == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Teams.BeforeDelete(ctx, actor, team) } @@ -166,6 +188,7 @@ func (e *ServiceHookExecutor) AfterDeleteOrganizationTeam(ctx context.Context, a if e == nil || e.config == nil || e.config.Teams == nil || e.config.Teams.AfterDelete == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.Teams.AfterDelete(ctx, actor, team) } @@ -173,6 +196,7 @@ func (e *ServiceHookExecutor) BeforeCreateOrganizationTeamMember(ctx context.Con if e == nil || e.config == nil || e.config.TeamMembers == nil || e.config.TeamMembers.BeforeCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.TeamMembers.BeforeCreate(ctx, actor, member) } @@ -180,6 +204,7 @@ func (e *ServiceHookExecutor) AfterCreateOrganizationTeamMember(ctx context.Cont if e == nil || e.config == nil || e.config.TeamMembers == nil || e.config.TeamMembers.AfterCreate == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.TeamMembers.AfterCreate(ctx, actor, member) } @@ -187,6 +212,7 @@ func (e *ServiceHookExecutor) BeforeDeleteOrganizationTeamMember(ctx context.Con if e == nil || e.config == nil || e.config.TeamMembers == nil || e.config.TeamMembers.BeforeDelete == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.TeamMembers.BeforeDelete(ctx, actor, member) } @@ -194,5 +220,6 @@ func (e *ServiceHookExecutor) AfterDeleteOrganizationTeamMember(ctx context.Cont if e == nil || e.config == nil || e.config.TeamMembers == nil || e.config.TeamMembers.AfterDelete == nil { return nil } + ctx = models.NewContextWithServiceRegistry(ctx, e.registry) return e.config.TeamMembers.AfterDelete(ctx, actor, member) } diff --git a/plugins/organizations/services/hooks_executor_test.go b/plugins/organizations/services/hooks_executor_test.go index a9cd367..a6ba422 100644 --- a/plugins/organizations/services/hooks_executor_test.go +++ b/plugins/organizations/services/hooks_executor_test.go @@ -12,7 +12,7 @@ import ( func TestServiceHookExecutor_NilConfigIsNoop(t *testing.T) { t.Parallel() - executor := NewServiceHookExecutor(nil) + executor := NewServiceHookExecutor(nil, nil) ctx := context.Background() actor := &models.Actor{ID: "user-1"} @@ -83,7 +83,7 @@ func TestServiceHookExecutor_OrganizationCreateHooks(t *testing.T) { return nil }, }, - }) + }, nil) ctx := context.Background() actor := &models.Actor{ID: "user-1"} @@ -114,7 +114,7 @@ func TestServiceHookExecutor_OrganizationCreateHookError(t *testing.T) { return someErr }, }, - }) + }, nil) err := executor.BeforeCreateOrganization(context.Background(), &models.Actor{ID: "user-1"}, &types.Organization{ID: "org-1"}) if !errors.Is(err, someErr) { @@ -161,7 +161,7 @@ func TestServiceHookExecutor_MemberUpdateDeleteHooks(t *testing.T) { return nil }, }, - }) + }, nil) ctx := context.Background() actor := &models.Actor{ID: "user-1"} @@ -184,3 +184,43 @@ func TestServiceHookExecutor_MemberUpdateDeleteHooks(t *testing.T) { t.Fatal("expected member update and delete hooks to be called") } } + +type testRegistry struct { + services map[string]any +} + +func (r *testRegistry) Register(name string, service any) {} + +func (r *testRegistry) Get(name string) any { + return r.services[name] +} + +func TestServiceHookExecutor_RegistryAccessibleInHook(t *testing.T) { + t.Parallel() + + var receivedService any + registry := &testRegistry{services: map[string]any{ + models.ServiceUser.String(): "mock-user-service", + }} + + executor := NewServiceHookExecutor(&types.OrganizationsServiceHooksConfig{ + Organizations: &types.OrganizationServiceHooksConfig{ + BeforeCreate: func(ctx context.Context, actor *models.Actor, organization *types.Organization) error { + ok, svc := models.GetServiceFromContext[string](ctx, models.ServiceUser) + if !ok { + return errors.New("service not found in context") + } + receivedService = svc + return nil + }, + }, + }, registry) + + err := executor.BeforeCreateOrganization(context.Background(), &models.Actor{ID: "user-1"}, &types.Organization{ID: "org-1"}) + if err != nil { + t.Fatalf("expected nil error, got %v", err) + } + if receivedService != "mock-user-service" { + t.Fatalf("expected 'mock-user-service', got %v", receivedService) + } +}