diff --git a/backend/internal/appconfig/app_config_middleware.go b/backend/internal/appconfig/app_config_middleware.go deleted file mode 100644 index 32619277..00000000 --- a/backend/internal/appconfig/app_config_middleware.go +++ /dev/null @@ -1,64 +0,0 @@ -package appconfig - -import ( - "context" - "errors" - "sync" - - "github.com/gin-gonic/gin" -) - -// appConfigCtxKey is the context key used to store the AppConfigResolver in the http.Request's context -type appConfigCtxKey struct{} - -type appConfigResolver func(ctx context.Context) (*AppConfigModel, error) - -// AppConfigMiddleware is a Gin middleware that makes the application configuration available to all downstream handlers through the request's context -type AppConfigMiddleware struct { - appConfigService *AppConfigService -} - -func NewAppConfigMiddleware(appConfigService *AppConfigService) *AppConfigMiddleware { - return &AppConfigMiddleware{ - appConfigService: appConfigService, - } -} - -// Add returns a Gin middleware that stores an AppConfigResolver in the http.Request's context -// The resolver loads the application configuration lazily on the first call and caches it for the duration of the request -func (m *AppConfigMiddleware) Add() gin.HandlerFunc { - return func(c *gin.Context) { - reqCtx := c.Request.Context() - - // Create a cache for each request in the middleware's scope, so it's unique per each request - var ( - once sync.Once - cfg *AppConfigModel - err error - ) - // Note: the resolver accepts a context argument, it doesn't use the request's own - // This can be used for example for tracing - resolver := appConfigResolver(func(ctx context.Context) (*AppConfigModel, error) { - once.Do(func() { - cfg, err = m.appConfigService.GetConfig(ctx) - }) - return cfg, err - }) - - // Store the resolver in the request's context - c.Request = c.Request.WithContext(context.WithValue(reqCtx, appConfigCtxKey{}, resolver)) - - c.Next() - } -} - -// FromCtx retrieves the app config from the context -func FromCtx(ctx context.Context) (*AppConfigModel, error) { - resolver, ok := ctx.Value(appConfigCtxKey{}).(appConfigResolver) - if !ok || resolver == nil { - // Indicates a development-time error - return nil, errors.New("middleware AppConfigMiddleware was not registered for the handler") - } - - return resolver(ctx) -} diff --git a/backend/internal/appconfig/testing_unit.go b/backend/internal/appconfig/testing_unit.go index 1e7718fb..3df1fa6c 100644 --- a/backend/internal/appconfig/testing_unit.go +++ b/backend/internal/appconfig/testing_unit.go @@ -3,10 +3,6 @@ // This file contains utils for unit tests and it's only built when the "unit" tag is set package appconfig -import ( - "context" -) - // NewTestAppConfigService is a function used by tests to create AppConfigService objects with pre-defined configuration values func NewTestAppConfigService(config *AppConfigModel) *AppConfigService { if config == nil { @@ -21,15 +17,11 @@ func NewTestAppConfigService(config *AppConfigModel) *AppConfigService { return service } -// NewTestContext returns a context that resolves the provided application configuration -func NewTestContext(ctx context.Context, config *AppConfigModel) context.Context { +// NewTestConfig returns an application configuration for use in tests, falling back to the default configuration when none is provided +func NewTestConfig(config *AppConfigModel) *AppConfigModel { if config == nil { config = getDefaultConfig() } - resolver := appConfigResolver(func(context.Context) (*AppConfigModel, error) { - return config, nil - }) - - return context.WithValue(ctx, appConfigCtxKey{}, resolver) + return config } diff --git a/backend/internal/bootstrap/router_bootstrap.go b/backend/internal/bootstrap/router_bootstrap.go index 3c6d0f76..11a71b96 100644 --- a/backend/internal/bootstrap/router_bootstrap.go +++ b/backend/internal/bootstrap/router_bootstrap.go @@ -23,7 +23,6 @@ import ( "gorm.io/gorm" "github.com/pocket-id/pocket-id/backend/frontend" - "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/controller" "github.com/pocket-id/pocket-id/backend/internal/middleware" @@ -141,14 +140,13 @@ func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, rateLimitServices } // Initialize middleware for specific routes - appConfigMiddleware := appconfig.NewAppConfigMiddleware(svc.appConfigService) authMiddleware := middleware.NewAuthMiddleware(svc.apiKeyModule, svc.userService, svc.jwtService) fileSizeLimitMiddleware := middleware.NewFileSizeLimitMiddleware() rateLimitMiddleware := middleware.NewRateLimitMiddleware(rateLimitServices) apiRateLimitMiddleware := rateLimitMiddleware.Add(middleware.RateLimitAPI) - apiGroup := r.Group("/api", appConfigMiddleware.Add(), apiRateLimitMiddleware) - baseGroup := r.Group("/", appConfigMiddleware.Add(), apiRateLimitMiddleware) + apiGroup := r.Group("/api", apiRateLimitMiddleware) + baseGroup := r.Group("/", apiRateLimitMiddleware) svc.apiKeyModule.RegisterRoutes(apiGroup, authMiddleware.WithAdminNotRequired().Add(), @@ -160,11 +158,11 @@ func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, rateLimitServices rateLimitMiddleware.Add(middleware.RateLimitWebauthnReauthenticate), ) controller.NewOidcController(apiGroup, authMiddleware, fileSizeLimitMiddleware, svc.oidcService) - controller.NewUserController(apiGroup, authMiddleware, rateLimitMiddleware, svc.userService, svc.oneTimeAccessService, svc.webauthnModule) + controller.NewUserController(apiGroup, authMiddleware, rateLimitMiddleware, svc.appConfigService, svc.userService, svc.oneTimeAccessService, svc.webauthnModule) controller.NewAppConfigController(apiGroup, authMiddleware, svc.appConfigService, svc.emailService, svc.ldapService) controller.NewAppImagesController(apiGroup, authMiddleware, svc.appImagesService) controller.NewAuditLogController(apiGroup, svc.auditLogService, authMiddleware) - controller.NewUserGroupController(apiGroup, authMiddleware, svc.userGroupService) + controller.NewUserGroupController(apiGroup, authMiddleware, svc.appConfigService, svc.userGroupService) svc.apiModule.RegisterRoutes(apiGroup, authMiddleware.Add()) controller.NewCustomClaimController(apiGroup, authMiddleware, svc.customClaimService) controller.NewVersionController(apiGroup, authMiddleware, svc.versionService) diff --git a/backend/internal/bootstrap/services_bootstrap.go b/backend/internal/bootstrap/services_bootstrap.go index a8f4f2f4..c677c05b 100644 --- a/backend/internal/bootstrap/services_bootstrap.go +++ b/backend/internal/bootstrap/services_bootstrap.go @@ -82,10 +82,11 @@ func initServices( svc.customClaimService = service.NewCustomClaimService(db) svc.webauthnModule, err = webauthn.New(webauthn.Dependencies{ - DB: db, - AppURL: common.EnvConfig.AppURL, - Signer: svc.jwtService, - AuditLog: svc.auditLogService, + DB: db, + AppURL: common.EnvConfig.AppURL, + Signer: svc.jwtService, + AuditLog: svc.auditLogService, + AppConfig: svc.appConfigService, }) if err != nil { return nil, fmt.Errorf("failed to create WebAuthn module: %w", err) @@ -136,6 +137,7 @@ func initServices( Signer: svc.jwtService, AuditLog: svc.auditLogService, UserCreator: svc.userService, + AppConfig: svc.appConfigService, }) svc.oneTimeAccessService = service.NewOneTimeAccessService(db, svc.userService, svc.jwtService, svc.auditLogService, svc.emailService) diff --git a/backend/internal/controller/app_config_controller.go b/backend/internal/controller/app_config_controller.go index 82e722c3..289cd765 100644 --- a/backend/internal/controller/app_config_controller.go +++ b/backend/internal/controller/app_config_controller.go @@ -53,7 +53,7 @@ type AppConfigController struct { // @Success 200 {array} dto.PublicAppConfigVariableDto // @Router /api/application-configuration [get] func (acc *AppConfigController) listAppConfigHandler(c *gin.Context) { - dbConfig, err := appconfig.FromCtx(c.Request.Context()) + dbConfig, err := acc.appConfigService.GetConfig(c.Request.Context()) if err != nil { _ = c.Error(err) return @@ -92,7 +92,7 @@ func (acc *AppConfigController) listAppConfigHandler(c *gin.Context) { // @Success 200 {array} dto.AppConfigVariableDto // @Router /api/application-configuration/all [get] func (acc *AppConfigController) listAllAppConfigHandler(c *gin.Context) { - dbConfig, err := appconfig.FromCtx(c.Request.Context()) + dbConfig, err := acc.appConfigService.GetConfig(c.Request.Context()) if err != nil { _ = c.Error(err) return @@ -146,7 +146,13 @@ func (acc *AppConfigController) updateAppConfigHandler(c *gin.Context) { // @Success 204 "No Content" // @Router /api/application-configuration/sync-ldap [post] func (acc *AppConfigController) syncLdapHandler(c *gin.Context) { - err := acc.ldapService.SyncAll(c.Request.Context()) + dbConfig, err := acc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(err) + return + } + + err = acc.ldapService.SyncAll(c.Request.Context(), dbConfig) if err != nil { _ = c.Error(err) return @@ -162,9 +168,15 @@ func (acc *AppConfigController) syncLdapHandler(c *gin.Context) { // @Success 204 "No Content" // @Router /api/application-configuration/test-email [post] func (acc *AppConfigController) testEmailHandler(c *gin.Context) { + dbConfig, err := acc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(err) + return + } + userID := c.GetString("userID") - err := acc.emailService.SendTestEmail(c.Request.Context(), userID) + err = acc.emailService.SendTestEmail(c.Request.Context(), dbConfig, userID) if err != nil { _ = c.Error(err) return diff --git a/backend/internal/controller/user_controller.go b/backend/internal/controller/user_controller.go index 5b13cd7a..57ea9095 100644 --- a/backend/internal/controller/user_controller.go +++ b/backend/internal/controller/user_controller.go @@ -23,8 +23,9 @@ const defaultOneTimeAccessTokenDuration = 15 * time.Minute // @Summary User management controller // @Description Initializes all user-related API endpoints // @Tags Users -func NewUserController(group *gin.RouterGroup, authMiddleware *middleware.AuthMiddleware, rateLimitMiddleware *middleware.RateLimitMiddleware, userService *service.UserService, oneTimeAccessService *service.OneTimeAccessService, webAuthnService *webauthn.Module) { +func NewUserController(group *gin.RouterGroup, authMiddleware *middleware.AuthMiddleware, rateLimitMiddleware *middleware.RateLimitMiddleware, appConfigService *appconfig.AppConfigService, userService *service.UserService, oneTimeAccessService *service.OneTimeAccessService, webAuthnService *webauthn.Module) { uc := UserController{ + appConfigService: appConfigService, userService: userService, oneTimeAccessService: oneTimeAccessService, webAuthnService: webAuthnService, @@ -62,6 +63,7 @@ func NewUserController(group *gin.RouterGroup, authMiddleware *middleware.AuthMi } type UserController struct { + appConfigService *appconfig.AppConfigService userService *service.UserService oneTimeAccessService *service.OneTimeAccessService webAuthnService *webauthn.Module @@ -207,7 +209,13 @@ func (uc *UserController) getCurrentUserHandler(c *gin.Context) { // @Success 204 "No Content" // @Router /api/users/{id} [delete] func (uc *UserController) deleteUserHandler(c *gin.Context) { - if err := uc.userService.DeleteUser(c.Request.Context(), c.Param("id"), false); err != nil { + dbConfig, err := uc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) + return + } + + if err := uc.userService.DeleteUser(c.Request.Context(), dbConfig, c.Param("id"), false); err != nil { _ = c.Error(err) return } @@ -248,13 +256,19 @@ func (uc *UserController) deleteUserWebauthnCredentialHandler(c *gin.Context) { // @Success 201 {object} dto.UserDto // @Router /api/users [post] func (uc *UserController) createUserHandler(c *gin.Context) { + dbConfig, err := uc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) + return + } + var input dto.UserCreateDto if err := dto.ShouldBindWithNormalizedJSON(c, &input); err != nil { _ = c.Error(err) return } - user, err := uc.userService.CreateUser(c.Request.Context(), input) + user, err := uc.userService.CreateUser(c.Request.Context(), dbConfig, input) if err != nil { _ = c.Error(err) return @@ -452,13 +466,19 @@ func (uc *UserController) createAdminOneTimeAccessTokenHandler(c *gin.Context) { // @Success 204 "No Content" // @Router /api/one-time-access-email [post] func (uc *UserController) RequestOneTimeAccessEmailAsUnauthenticatedUserHandler(c *gin.Context) { + dbConfig, err := uc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) + return + } + var input dto.OneTimeAccessEmailAsUnauthenticatedUserDto if err := dto.ShouldBindWithNormalizedJSON(c, &input); err != nil { _ = c.Error(err) return } - deviceToken, err := uc.oneTimeAccessService.RequestOneTimeAccessEmailAsUnauthenticatedUser(c.Request.Context(), input.Email, input.RedirectPath) + deviceToken, err := uc.oneTimeAccessService.RequestOneTimeAccessEmailAsUnauthenticatedUser(c.Request.Context(), dbConfig, input.Email, input.RedirectPath) if err != nil { _ = c.Error(err) return @@ -479,6 +499,12 @@ func (uc *UserController) RequestOneTimeAccessEmailAsUnauthenticatedUserHandler( // @Success 204 "No Content" // @Router /api/users/{id}/one-time-access-email [post] func (uc *UserController) RequestOneTimeAccessEmailAsAdminHandler(c *gin.Context) { + dbConfig, err := uc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) + return + } + var input dto.OneTimeAccessEmailAsAdminDto if err := c.ShouldBindJSON(&input); err != nil { _ = c.Error(err) @@ -491,7 +517,7 @@ func (uc *UserController) RequestOneTimeAccessEmailAsAdminHandler(c *gin.Context if ttl <= 0 { ttl = defaultOneTimeAccessTokenDuration } - err := uc.oneTimeAccessService.RequestOneTimeAccessEmailAsAdmin(c.Request.Context(), userID, ttl) + err = uc.oneTimeAccessService.RequestOneTimeAccessEmailAsAdmin(c.Request.Context(), dbConfig, userID, ttl) if err != nil { _ = c.Error(err) return @@ -508,7 +534,7 @@ func (uc *UserController) RequestOneTimeAccessEmailAsAdminHandler(c *gin.Context // @Success 200 {object} dto.UserDto // @Router /api/one-time-access-token/{token} [post] func (uc *UserController) exchangeOneTimeAccessTokenHandler(c *gin.Context) { - cfg, err := appconfig.FromCtx(c.Request.Context()) + cfg, err := uc.appConfigService.GetConfig(c.Request.Context()) if err != nil { _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) return @@ -522,7 +548,7 @@ func (uc *UserController) exchangeOneTimeAccessTokenHandler(c *gin.Context) { } deviceToken, _ := c.Cookie(cookie.DeviceTokenCookieName) - user, token, err := uc.oneTimeAccessService.ExchangeOneTimeAccessToken(c.Request.Context(), loginCode, deviceToken, c.ClientIP(), c.Request.UserAgent()) + user, token, err := uc.oneTimeAccessService.ExchangeOneTimeAccessToken(c.Request.Context(), cfg, loginCode, deviceToken, c.ClientIP(), c.Request.UserAgent()) if err != nil { _ = c.Error(err) return @@ -573,6 +599,12 @@ func (uc *UserController) updateUserGroups(c *gin.Context) { // updateUser is an internal helper method, not exposed as an API endpoint func (uc *UserController) updateUser(c *gin.Context, updateOwnUser bool) { + dbConfig, err := uc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) + return + } + var input dto.UserCreateDto if err := dto.ShouldBindWithNormalizedJSON(c, &input); err != nil { _ = c.Error(err) @@ -586,7 +618,7 @@ func (uc *UserController) updateUser(c *gin.Context, updateOwnUser bool) { userID = c.Param("id") } - user, err := uc.userService.UpdateUser(c.Request.Context(), userID, input, updateOwnUser, false) + user, err := uc.userService.UpdateUser(c.Request.Context(), dbConfig, userID, input, updateOwnUser, false) if err != nil { _ = c.Error(err) return @@ -646,9 +678,15 @@ func (uc *UserController) resetCurrentUserProfilePictureHandler(c *gin.Context) // @Success 204 "No Content" // @Router /api/users/me/send-email-verification [post] func (uc *UserController) sendEmailVerificationHandler(c *gin.Context) { + dbConfig, err := uc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) + return + } + userID := c.GetString("userID") - if err := uc.userService.SendEmailVerification(c.Request.Context(), userID); err != nil { + if err := uc.userService.SendEmailVerification(c.Request.Context(), dbConfig, userID); err != nil { _ = c.Error(err) return } diff --git a/backend/internal/controller/user_group_controller.go b/backend/internal/controller/user_group_controller.go index a0567752..5ecfd319 100644 --- a/backend/internal/controller/user_group_controller.go +++ b/backend/internal/controller/user_group_controller.go @@ -1,9 +1,11 @@ package controller import ( + "fmt" "net/http" "github.com/gin-gonic/gin" + "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/middleware" "github.com/pocket-id/pocket-id/backend/internal/service" @@ -14,8 +16,9 @@ import ( // @Summary User group management controller // @Description Initializes all user group-related API endpoints // @Tags User Groups -func NewUserGroupController(group *gin.RouterGroup, authMiddleware *middleware.AuthMiddleware, userGroupService *service.UserGroupService) { +func NewUserGroupController(group *gin.RouterGroup, authMiddleware *middleware.AuthMiddleware, appConfigService *appconfig.AppConfigService, userGroupService *service.UserGroupService) { ugc := UserGroupController{ + appConfigService: appConfigService, UserGroupService: userGroupService, } @@ -33,6 +36,7 @@ func NewUserGroupController(group *gin.RouterGroup, authMiddleware *middleware.A } type UserGroupController struct { + appConfigService *appconfig.AppConfigService UserGroupService *service.UserGroupService } @@ -146,13 +150,19 @@ func (ugc *UserGroupController) create(c *gin.Context) { // @Success 200 {object} dto.UserGroupDto "Updated user group" // @Router /api/user-groups/{id} [put] func (ugc *UserGroupController) update(c *gin.Context) { + dbConfig, err := ugc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) + return + } + var input dto.UserGroupCreateDto if err := dto.ShouldBindWithNormalizedJSON(c, &input); err != nil { _ = c.Error(err) return } - group, err := ugc.UserGroupService.Update(c.Request.Context(), c.Param("id"), input) + group, err := ugc.UserGroupService.Update(c.Request.Context(), dbConfig, c.Param("id"), input) if err != nil { _ = c.Error(err) return @@ -177,7 +187,13 @@ func (ugc *UserGroupController) update(c *gin.Context) { // @Success 204 "No Content" // @Router /api/user-groups/{id} [delete] func (ugc *UserGroupController) delete(c *gin.Context) { - if err := ugc.UserGroupService.Delete(c.Request.Context(), c.Param("id")); err != nil { + dbConfig, err := ugc.appConfigService.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) + return + } + + if err := ugc.UserGroupService.Delete(c.Request.Context(), dbConfig, c.Param("id")); err != nil { _ = c.Error(err) return } diff --git a/backend/internal/job/api_key_expiry_job.go b/backend/internal/job/api_key_expiry_job.go index 74e63fb4..4717f0a2 100644 --- a/backend/internal/job/api_key_expiry_job.go +++ b/backend/internal/job/api_key_expiry_job.go @@ -51,7 +51,7 @@ func (j *ApiKeyEmailJobs) checkAndNotifyExpiringApiKeys(ctx context.Context) err continue } - err = service.SendEmailWithConfig(ctx, j.emailService, dbConfig, email.Address{ + err = service.SendEmail(ctx, j.emailService, dbConfig, email.Address{ Name: key.User.FullName(), Email: *key.User.Email, }, service.ApiKeyExpiringSoonTemplate, &service.ApiKeyExpiringSoonTemplateData{ diff --git a/backend/internal/job/ldap_job.go b/backend/internal/job/ldap_job.go index 1f68d47a..4c450bcb 100644 --- a/backend/internal/job/ldap_job.go +++ b/backend/internal/job/ldap_job.go @@ -31,5 +31,5 @@ func (j *LdapJobs) syncLdap(ctx context.Context) error { return nil } - return j.ldapService.SyncAllWithConfig(ctx, dbConfig) + return j.ldapService.SyncAll(ctx, dbConfig) } diff --git a/backend/internal/service/audit_log_service.go b/backend/internal/service/audit_log_service.go index 290aedce..4d8bf905 100644 --- a/backend/internal/service/audit_log_service.go +++ b/backend/internal/service/audit_log_service.go @@ -111,7 +111,7 @@ func (s *AuditLogService) CreateNewSignInWithEmail(ctx context.Context, ipAddres return } - innerErr = SendEmail(innerCtx, s.emailService, email.Address{ + innerErr = SendEmail(innerCtx, s.emailService, dbConfig, email.Address{ Name: user.FullName(), Email: *user.Email, }, NewLoginTemplate, &NewLoginTemplateData{ diff --git a/backend/internal/service/e2etest_service.go b/backend/internal/service/e2etest_service.go index 16b691c7..21572adc 100644 --- a/backend/internal/service/e2etest_service.go +++ b/backend/internal/service/e2etest_service.go @@ -663,7 +663,11 @@ func (s *TestService) ResetLock(ctx context.Context) error { // SyncLdap triggers an LDAP synchronization func (s *TestService) SyncLdap(ctx context.Context) error { - return s.ldapService.SyncAll(ctx) + dbConfig, err := s.appConfigService.GetConfig(ctx) + if err != nil { + return fmt.Errorf("error loading app configuration: %w", err) + } + return s.ldapService.SyncAll(ctx, dbConfig) } // SetLdapTestConfig updates the LDAP configuration used by the end-to-end test server diff --git a/backend/internal/service/email_service.go b/backend/internal/service/email_service.go index b73b6531..98ac360a 100644 --- a/backend/internal/service/email_service.go +++ b/backend/internal/service/email_service.go @@ -43,7 +43,7 @@ func NewEmailService(db *gorm.DB) (*EmailService, error) { }, nil } -func (srv *EmailService) SendTestEmail(ctx context.Context, recipientUserId string) error { +func (srv *EmailService) SendTestEmail(ctx context.Context, dbConfig *appconfig.AppConfigModel, recipientUserId string) error { var user model.User err := srv.db. WithContext(ctx). @@ -57,24 +57,15 @@ func (srv *EmailService) SendTestEmail(ctx context.Context, recipientUserId stri return &common.UserEmailNotSetError{} } - return SendEmail(ctx, srv, + return SendEmail(ctx, srv, dbConfig, email.Address{ Email: *user.Email, Name: user.FullName(), }, TestTemplate, nil) } -func SendEmail[V any](ctx context.Context, srv *EmailService, toEmail email.Address, template email.Template[V], tData *V) error { - dbConfig, err := appconfig.FromCtx(ctx) - if err != nil { - return fmt.Errorf("error loading app configuration: %w", err) - } - - return SendEmailWithConfig(ctx, srv, dbConfig, toEmail, template, tData) -} - -// SendEmailWithConfig sends an email with an explicitly loaded configuration for call chains that do not originate from an HTTP request -func SendEmailWithConfig[V any](ctx context.Context, srv *EmailService, dbConfig *appconfig.AppConfigModel, toEmail email.Address, template email.Template[V], tData *V) error { +// SendEmail sends an email using the provided application configuration +func SendEmail[V any](ctx context.Context, srv *EmailService, dbConfig *appconfig.AppConfigModel, toEmail email.Address, template email.Template[V], tData *V) error { data := &email.TemplateData[V]{ AppName: dbConfig.AppName.String(), diff --git a/backend/internal/service/ldap_service.go b/backend/internal/service/ldap_service.go index 2bf917ec..34f6a3f4 100644 --- a/backend/internal/service/ldap_service.go +++ b/backend/internal/service/ldap_service.go @@ -103,20 +103,8 @@ func (s *LdapService) createClient(dbConfig *appconfig.AppConfigModel) (ldapClie return client, nil } -func (s *LdapService) SyncAll(ctx context.Context) error { - dbConfig, err := appconfig.FromCtx(ctx) - if err != nil { - return fmt.Errorf("error loading app configuration: %w", err) - } - return s.syncAll(ctx, dbConfig) -} - -// SyncAllWithConfig synchronizes LDAP with an explicitly loaded configuration for call chains that do not originate from an HTTP request -func (s *LdapService) SyncAllWithConfig(ctx context.Context, dbConfig *appconfig.AppConfigModel) error { - return s.syncAll(ctx, dbConfig) -} - -func (s *LdapService) syncAll(ctx context.Context, dbConfig *appconfig.AppConfigModel) error { +// SyncAll synchronizes LDAP using the provided application configuration +func (s *LdapService) SyncAll(ctx context.Context, dbConfig *appconfig.AppConfigModel) error { // Setup LDAP connection client, err := s.clientFactory(dbConfig) if err != nil { @@ -144,7 +132,7 @@ func (s *LdapService) syncAll(ctx context.Context, dbConfig *appconfig.AppConfig } // Reconcile groups - err = s.reconcileGroups(ctx, tx, desiredState.groups, desiredState.groupIDs) + err = s.reconcileGroups(ctx, tx, desiredState.groups, desiredState.groupIDs, dbConfig) if err != nil { return fmt.Errorf("failed to sync groups: %w", err) } @@ -426,7 +414,7 @@ func (s *LdapService) resolveGroupMemberUsername(ctx context.Context, client lda return norm.NFC.String(username) } -func (s *LdapService) reconcileGroups(ctx context.Context, tx *gorm.DB, desiredGroups []ldapDesiredGroup, ldapGroupIDs map[string]struct{}) error { +func (s *LdapService) reconcileGroups(ctx context.Context, tx *gorm.DB, desiredGroups []ldapDesiredGroup, ldapGroupIDs map[string]struct{}, dbConfig *appconfig.AppConfigModel) error { // Load the current LDAP-managed state from the database ldapGroupsInDB, ldapGroupsByID, err := s.loadLDAPGroupsInDB(ctx, tx) if err != nil { @@ -466,7 +454,7 @@ func (s *LdapService) reconcileGroups(ctx context.Context, tx *gorm.DB, desiredG continue } - _, err = s.groupService.updateInternal(ctx, databaseGroup.ID, desiredGroup.input, true, tx) + _, err = s.groupService.updateInternal(ctx, databaseGroup.ID, desiredGroup.input, true, tx, dbConfig) if err != nil { return fmt.Errorf("failed to update group '%s': %w", desiredGroup.input.Name, err) } @@ -583,7 +571,7 @@ func (s *LdapService) reconcileUsers(ctx context.Context, tx *gorm.DB, desiredUs continue } - err = s.userService.deleteUserInternal(ctx, tx, user.ID, true) + err = s.userService.deleteUserInternal(ctx, tx, user.ID, true, dbConfig) if err != nil { if _, ok := errors.AsType[*common.LdapUserUpdateError](err); ok { return nil, nil, fmt.Errorf("failed to delete user %s: LDAP user must be disabled before deletion", user.Username) diff --git a/backend/internal/service/ldap_service_test.go b/backend/internal/service/ldap_service_test.go index 540f25ab..6f935fc9 100644 --- a/backend/internal/service/ldap_service_test.go +++ b/backend/internal/service/ldap_service_test.go @@ -110,7 +110,7 @@ func TestLdapServiceSyncAllReconcilesUsersAndGroups(t *testing.T) { LdapID: &oldGroupLdapID, }).Error) - err := service.SyncAllWithConfig(t.Context(), defaultTestLDAPAppConfig()) + err := service.SyncAll(t.Context(), defaultTestLDAPAppConfig()) require.NoError(t, err) var alice model.User @@ -177,7 +177,7 @@ func TestLdapServiceSyncAllMapsPosixGroupMemberUid(t *testing.T) { ), )) - err := service.SyncAllWithConfig(t.Context(), appCfg) + err := service.SyncAll(t.Context(), appCfg) require.NoError(t, err) var group model.UserGroup @@ -220,7 +220,7 @@ func TestLdapServiceSyncAllHandlesDuplicateLDAPIDsInSingleRun(t *testing.T) { ), )) - err := service.SyncAllWithConfig(t.Context(), defaultTestLDAPAppConfig()) + err := service.SyncAll(t.Context(), defaultTestLDAPAppConfig()) require.NoError(t, err) var users []model.User @@ -292,7 +292,7 @@ func TestLdapServiceSyncAllSetsAdminFromGroupMembership(t *testing.T) { ldapSearchResult(tt.groupEntry), )) - err := service.SyncAllWithConfig(t.Context(), tt.appConfig) + err := service.SyncAll(t.Context(), tt.appConfig) require.NoError(t, err) var user model.User diff --git a/backend/internal/service/one_time_access_service.go b/backend/internal/service/one_time_access_service.go index c95f84a3..af561ee3 100644 --- a/backend/internal/service/one_time_access_service.go +++ b/backend/internal/service/one_time_access_service.go @@ -37,32 +37,22 @@ func NewOneTimeAccessService(db *gorm.DB, userService *UserService, jwtService * } } -func (s *OneTimeAccessService) RequestOneTimeAccessEmailAsAdmin(ctx context.Context, userID string, ttl time.Duration) error { - dbConfig, err := appconfig.FromCtx(ctx) - if err != nil { - return fmt.Errorf("error loading app configuration: %w", err) - } - +func (s *OneTimeAccessService) RequestOneTimeAccessEmailAsAdmin(ctx context.Context, dbConfig *appconfig.AppConfigModel, userID string, ttl time.Duration) error { if !dbConfig.EmailOneTimeAccessAsAdminEnabled.IsTrue() { return &common.OneTimeAccessDisabledError{} } - _, err = s.requestOneTimeAccessEmailInternal(ctx, userID, "", ttl, false) + _, err := s.requestOneTimeAccessEmailInternal(ctx, userID, "", ttl, false, dbConfig) return err } -func (s *OneTimeAccessService) RequestOneTimeAccessEmailAsUnauthenticatedUser(ctx context.Context, userID, redirectPath string) (string, error) { - dbConfig, err := appconfig.FromCtx(ctx) - if err != nil { - return "", fmt.Errorf("error loading app configuration: %w", err) - } - +func (s *OneTimeAccessService) RequestOneTimeAccessEmailAsUnauthenticatedUser(ctx context.Context, dbConfig *appconfig.AppConfigModel, userID, redirectPath string) (string, error) { if !dbConfig.EmailOneTimeAccessAsUnauthenticatedEnabled.IsTrue() { return "", &common.OneTimeAccessDisabledError{} } var userId string - err = s.db.Model(&model.User{}).Select("id").Where("email = ?", userID).First(&userId).Error + err := s.db.Model(&model.User{}).Select("id").Where("email = ?", userID).First(&userId).Error if errors.Is(err, gorm.ErrRecordNotFound) { // Do not return error if user not found to prevent email enumeration return "", nil @@ -70,7 +60,7 @@ func (s *OneTimeAccessService) RequestOneTimeAccessEmailAsUnauthenticatedUser(ct return "", err } - deviceToken, err := s.requestOneTimeAccessEmailInternal(ctx, userId, redirectPath, 15*time.Minute, true) + deviceToken, err := s.requestOneTimeAccessEmailInternal(ctx, userId, redirectPath, 15*time.Minute, true, dbConfig) if err != nil { return "", err } else if deviceToken == nil { @@ -80,7 +70,7 @@ func (s *OneTimeAccessService) RequestOneTimeAccessEmailAsUnauthenticatedUser(ct return *deviceToken, nil } -func (s *OneTimeAccessService) requestOneTimeAccessEmailInternal(ctx context.Context, userID, redirectPath string, ttl time.Duration, withDeviceToken bool) (*string, error) { +func (s *OneTimeAccessService) requestOneTimeAccessEmailInternal(ctx context.Context, userID, redirectPath string, ttl time.Duration, withDeviceToken bool, dbConfig *appconfig.AppConfigModel) (*string, error) { tx := s.db.Begin() defer func() { tx.Rollback() @@ -118,7 +108,7 @@ func (s *OneTimeAccessService) requestOneTimeAccessEmailInternal(ctx context.Con linkWithCode = linkWithCode + "?redirect=" + encodedRedirectPath } - errInternal := SendEmail(innerCtx, s.emailService, email.Address{ + errInternal := SendEmail(innerCtx, s.emailService, dbConfig, email.Address{ Name: user.FullName(), Email: *user.Email, }, OneTimeAccessTemplate, &OneTimeAccessTemplateData{ @@ -179,19 +169,14 @@ func (s *OneTimeAccessService) createOneTimeAccessTokenInternal(ctx context.Cont return oneTimeAccessToken.Token, oneTimeAccessToken.DeviceToken, nil } -func (s *OneTimeAccessService) ExchangeOneTimeAccessToken(ctx context.Context, token, deviceToken, ipAddress, userAgent string) (model.User, string, error) { - dbConfig, err := appconfig.FromCtx(ctx) - if err != nil { - return model.User{}, "", fmt.Errorf("error loading app configuration: %w", err) - } - +func (s *OneTimeAccessService) ExchangeOneTimeAccessToken(ctx context.Context, dbConfig *appconfig.AppConfigModel, token, deviceToken, ipAddress, userAgent string) (model.User, string, error) { tx := s.db.Begin() defer func() { tx.Rollback() }() var oneTimeAccessToken model.OneTimeAccessToken - err = tx. + err := tx. WithContext(ctx). Where("token = ? AND expires_at > ?", token, datatype.DateTime(time.Now())). Preload("User"). diff --git a/backend/internal/service/one_time_access_service_test.go b/backend/internal/service/one_time_access_service_test.go index 3f819c19..dd82fd71 100644 --- a/backend/internal/service/one_time_access_service_test.go +++ b/backend/internal/service/one_time_access_service_test.go @@ -36,8 +36,8 @@ func TestExchangeOneTimeAccessTokenRejectsDisabledUser(t *testing.T) { } require.NoError(t, db.Create(&loginCode).Error) - ctx := appconfig.NewTestContext(t.Context(), nil) - exchangedUser, accessToken, err := oneTimeAccessService.ExchangeOneTimeAccessToken(ctx, loginCode.Token, "", "", "") + dbConfig := appconfig.NewTestConfig(nil) + exchangedUser, accessToken, err := oneTimeAccessService.ExchangeOneTimeAccessToken(t.Context(), dbConfig, loginCode.Token, "", "", "") var userDisabledErr *common.UserDisabledError require.ErrorAs(t, err, &userDisabledErr) diff --git a/backend/internal/service/user_group_service.go b/backend/internal/service/user_group_service.go index de8a8876..fe942250 100644 --- a/backend/internal/service/user_group_service.go +++ b/backend/internal/service/user_group_service.go @@ -3,7 +3,6 @@ package service import ( "context" "errors" - "fmt" "time" "github.com/pocket-id/pocket-id/backend/internal/appconfig" @@ -63,19 +62,14 @@ func (s *UserGroupService) getInternal(ctx context.Context, id string, tx *gorm. return group, err } -func (s *UserGroupService) Delete(ctx context.Context, id string) error { - cfg, err := appconfig.FromCtx(ctx) - if err != nil { - return fmt.Errorf("error loading app configuration: %w", err) - } - +func (s *UserGroupService) Delete(ctx context.Context, cfg *appconfig.AppConfigModel, id string) error { tx := s.db.Begin() defer func() { tx.Rollback() }() var group model.UserGroup - err = tx. + err := tx. WithContext(ctx). Where("id = ?", id). First(&group). @@ -141,13 +135,13 @@ func (s *UserGroupService) createInternal(ctx context.Context, input dto.UserGro return group, nil } -func (s *UserGroupService) Update(ctx context.Context, id string, input dto.UserGroupCreateDto) (group model.UserGroup, err error) { +func (s *UserGroupService) Update(ctx context.Context, cfg *appconfig.AppConfigModel, id string, input dto.UserGroupCreateDto) (group model.UserGroup, err error) { tx := s.db.Begin() defer func() { tx.Rollback() }() - group, err = s.updateInternal(ctx, id, input, false, tx) + group, err = s.updateInternal(ctx, id, input, false, tx, cfg) if err != nil { return model.UserGroup{}, err } @@ -160,7 +154,7 @@ func (s *UserGroupService) Update(ctx context.Context, id string, input dto.User return group, nil } -func (s *UserGroupService) updateInternal(ctx context.Context, id string, input dto.UserGroupCreateDto, isLdapSync bool, tx *gorm.DB) (group model.UserGroup, err error) { +func (s *UserGroupService) updateInternal(ctx context.Context, id string, input dto.UserGroupCreateDto, isLdapSync bool, tx *gorm.DB, cfg *appconfig.AppConfigModel) (group model.UserGroup, err error) { group, err = s.getInternal(ctx, id, tx) if err != nil { return model.UserGroup{}, err @@ -168,10 +162,6 @@ func (s *UserGroupService) updateInternal(ctx context.Context, id string, input // Disallow updating the group if it is an LDAP group and LDAP is enabled if !isLdapSync && group.LdapID != nil { - cfg, err := appconfig.FromCtx(ctx) - if err != nil { - return model.UserGroup{}, fmt.Errorf("error loading app configuration: %w", err) - } if cfg.LdapEnabled.IsTrue() { return model.UserGroup{}, &common.LdapUserGroupUpdateError{} } diff --git a/backend/internal/service/user_service.go b/backend/internal/service/user_service.go index 72c0b847..34a1ee07 100644 --- a/backend/internal/service/user_service.go +++ b/backend/internal/service/user_service.go @@ -184,9 +184,9 @@ func (s *UserService) UpdateProfilePicture(ctx context.Context, userID string, f return nil } -func (s *UserService) DeleteUser(ctx context.Context, userID string, allowLdapDelete bool) error { +func (s *UserService) DeleteUser(ctx context.Context, dbConfig *appconfig.AppConfigModel, userID string, allowLdapDelete bool) error { err := s.db.Transaction(func(tx *gorm.DB) error { - return s.deleteUserInternal(ctx, tx, userID, allowLdapDelete) + return s.deleteUserInternal(ctx, tx, userID, allowLdapDelete, dbConfig) }) if err != nil { return fmt.Errorf("failed to delete user '%s': %w", userID, err) @@ -202,7 +202,7 @@ func (s *UserService) DeleteUser(ctx context.Context, userID string, allowLdapDe return nil } -func (s *UserService) deleteUserInternal(ctx context.Context, tx *gorm.DB, userID string, allowLdapDelete bool) error { +func (s *UserService) deleteUserInternal(ctx context.Context, tx *gorm.DB, userID string, allowLdapDelete bool, cfg *appconfig.AppConfigModel) error { var user model.User err := tx. WithContext(ctx). @@ -216,10 +216,6 @@ func (s *UserService) deleteUserInternal(ctx context.Context, tx *gorm.DB, userI // Disallow deleting the user if it is an LDAP user, LDAP is enabled, and the user is not disabled if !allowLdapDelete && !user.Disabled && user.LdapID != nil { - cfg, err := appconfig.FromCtx(ctx) - if err != nil { - return fmt.Errorf("error loading app configuration: %w", err) - } if cfg.LdapEnabled.IsTrue() { return &common.LdapUserUpdateError{} } @@ -237,13 +233,13 @@ func (s *UserService) deleteUserInternal(ctx context.Context, tx *gorm.DB, userI return nil } -func (s *UserService) CreateUser(ctx context.Context, input dto.UserCreateDto) (model.User, error) { +func (s *UserService) CreateUser(ctx context.Context, dbConfig *appconfig.AppConfigModel, input dto.UserCreateDto) (model.User, error) { tx := s.db.Begin() defer func() { tx.Rollback() }() - user, err := s.CreateUserInternal(ctx, input, false, tx) + user, err := s.CreateUserInternal(ctx, dbConfig, input, false, tx) if err != nil { return model.User{}, err } @@ -256,12 +252,8 @@ func (s *UserService) CreateUser(ctx context.Context, input dto.UserCreateDto) ( return user, nil } -func (s *UserService) CreateUserInternal(ctx context.Context, input dto.UserCreateDto, isLdapSync bool, tx *gorm.DB) (model.User, error) { - cfg, err := appconfig.FromCtx(ctx) - if err != nil { - return model.User{}, fmt.Errorf("error loading app configuration: %w", err) - } - return s.createUserInternal(ctx, input, isLdapSync, tx, cfg) +func (s *UserService) CreateUserInternal(ctx context.Context, dbConfig *appconfig.AppConfigModel, input dto.UserCreateDto, isLdapSync bool, tx *gorm.DB) (model.User, error) { + return s.createUserInternal(ctx, input, isLdapSync, tx, dbConfig) } func (s *UserService) createUserInternal(ctx context.Context, input dto.UserCreateDto, isLdapSync bool, tx *gorm.DB, cfg *appconfig.AppConfigModel) (model.User, error) { @@ -425,12 +417,7 @@ func (s *UserService) applyDefaultCustomClaims(ctx context.Context, user *model. return nil } -func (s *UserService) UpdateUser(ctx context.Context, userID string, updatedUser dto.UserCreateDto, updateOwnUser bool, isLdapSync bool) (model.User, error) { - cfg, err := appconfig.FromCtx(ctx) - if err != nil { - return model.User{}, fmt.Errorf("error loading app configuration: %w", err) - } - +func (s *UserService) UpdateUser(ctx context.Context, cfg *appconfig.AppConfigModel, userID string, updatedUser dto.UserCreateDto, updateOwnUser bool, isLdapSync bool) (model.User, error) { tx := s.db.Begin() defer func() { tx.Rollback() @@ -652,7 +639,7 @@ func (s *UserService) disableUserInternal(ctx context.Context, tx *gorm.DB, user return nil } -func (s *UserService) SendEmailVerification(ctx context.Context, userID string) error { +func (s *UserService) SendEmailVerification(ctx context.Context, dbConfig *appconfig.AppConfigModel, userID string) error { user, err := s.GetUser(ctx, userID) if err != nil { return err @@ -679,7 +666,7 @@ func (s *UserService) SendEmailVerification(ctx context.Context, userID string) return err } - return SendEmail(ctx, s.emailService, email.Address{ + return SendEmail(ctx, s.emailService, dbConfig, email.Address{ Name: user.FullName(), Email: *user.Email, }, EmailVerificationTemplate, &EmailVerificationTemplateData{ diff --git a/backend/internal/service/user_service_test.go b/backend/internal/service/user_service_test.go index a7d584a7..3b885d69 100644 --- a/backend/internal/service/user_service_test.go +++ b/backend/internal/service/user_service_test.go @@ -37,7 +37,6 @@ func newTestUserService(t *testing.T) (*UserService, *UserGroupService) { func TestCreateUserBumpsGroupUpdatedAt(t *testing.T) { config := &appconfig.AppConfigModel{RequireUserEmail: "false"} - ctx := appconfig.NewTestContext(t.Context(), config) userService, groupService := newTestUserService(t) group, err := groupService.Create(t.Context(), dto.UserGroupCreateDto{ @@ -50,7 +49,7 @@ func TestCreateUserBumpsGroupUpdatedAt(t *testing.T) { // Create a user that is a member of the group // This mirrors signing up via an invite link that adds the user to a group email := "member@example.com" - _, err = userService.CreateUser(ctx, dto.UserCreateDto{ + _, err = userService.CreateUser(t.Context(), config, dto.UserCreateDto{ Username: "member", Email: &email, FirstName: "Group", @@ -69,7 +68,6 @@ func TestCreateUserBumpsGroupUpdatedAt(t *testing.T) { func TestCreateUserBumpsDefaultGroupUpdatedAt(t *testing.T) { config := &appconfig.AppConfigModel{RequireUserEmail: "false"} - ctx := appconfig.NewTestContext(t.Context(), config) userService, groupService := newTestUserService(t) group, err := groupService.Create(t.Context(), dto.UserGroupCreateDto{ @@ -86,7 +84,7 @@ func TestCreateUserBumpsDefaultGroupUpdatedAt(t *testing.T) { // Create a user without explicit group IDs, so the default groups apply email := "default@example.com" - _, err = userService.CreateUser(ctx, dto.UserCreateDto{ + _, err = userService.CreateUser(t.Context(), config, dto.UserCreateDto{ Username: "defaultmember", Email: &email, FirstName: "Default", diff --git a/backend/internal/usersignup/handler.go b/backend/internal/usersignup/handler.go index 1dda51c0..c8d96ef2 100644 --- a/backend/internal/usersignup/handler.go +++ b/backend/internal/usersignup/handler.go @@ -7,7 +7,6 @@ import ( "github.com/gin-gonic/gin" - "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/utils" @@ -17,11 +16,12 @@ import ( const defaultSignupTokenDuration = time.Hour type handler struct { - service *Service + service *Service + appConfig AppConfigResolver } -func newHandler(service *Service) *handler { - return &handler{service: service} +func newHandler(service *Service, appConfig AppConfigResolver) *handler { + return &handler{service: service, appConfig: appConfig} } func (h *handler) checkInitialAdminSetupAvailable(c *gin.Context) { @@ -49,7 +49,7 @@ func (h *handler) checkInitialAdminSetupAvailable(c *gin.Context) { // @Success 200 {object} dto.UserDto // @Router /api/signup/setup [post] func (h *handler) signUpInitialAdmin(c *gin.Context) { - config, err := appconfig.FromCtx(c.Request.Context()) + config, err := h.appConfig.GetConfig(c.Request.Context()) if err != nil { _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) return @@ -61,7 +61,7 @@ func (h *handler) signUpInitialAdmin(c *gin.Context) { return } - user, token, err := h.service.SignUpInitialAdmin(c.Request.Context(), input) + user, token, err := h.service.SignUpInitialAdmin(c.Request.Context(), config, input) if err != nil { _ = c.Error(err) return @@ -176,7 +176,7 @@ func (h *handler) deleteSignupToken(c *gin.Context) { // @Success 201 {object} dto.UserDto // @Router /api/signup [post] func (h *handler) signup(c *gin.Context) { - config, err := appconfig.FromCtx(c.Request.Context()) + config, err := h.appConfig.GetConfig(c.Request.Context()) if err != nil { _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) return @@ -191,7 +191,7 @@ func (h *handler) signup(c *gin.Context) { ipAddress := c.ClientIP() userAgent := c.GetHeader("User-Agent") - user, accessToken, err := h.service.SignUp(c.Request.Context(), input, ipAddress, userAgent) + user, accessToken, err := h.service.SignUp(c.Request.Context(), config, input, ipAddress, userAgent) if err != nil { _ = c.Error(err) return diff --git a/backend/internal/usersignup/module.go b/backend/internal/usersignup/module.go index a4fc57bd..8a7acacf 100644 --- a/backend/internal/usersignup/module.go +++ b/backend/internal/usersignup/module.go @@ -7,6 +7,7 @@ import ( "github.com/gin-gonic/gin" "gorm.io/gorm" + "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/model" ) @@ -20,7 +21,12 @@ type AuditLogger interface { } type UserCreator interface { - CreateUserInternal(ctx context.Context, input dto.UserCreateDto, isLdapSync bool, tx *gorm.DB) (model.User, error) + CreateUserInternal(ctx context.Context, dbConfig *appconfig.AppConfigModel, input dto.UserCreateDto, isLdapSync bool, tx *gorm.DB) (model.User, error) +} + +// AppConfigResolver loads the current application configuration, so handlers can pass it explicitly to the service methods that need it +type AppConfigResolver interface { + GetConfig(ctx context.Context) (*appconfig.AppConfigModel, error) } type Dependencies struct { @@ -29,6 +35,7 @@ type Dependencies struct { Signer TokenService AuditLog AuditLogger UserCreator UserCreator + AppConfig AppConfigResolver } type Module struct { @@ -40,7 +47,7 @@ func New(deps Dependencies) *Module { service := newService(deps) return &Module{ service: service, - handler: newHandler(service), + handler: newHandler(service, deps.AppConfig), } } diff --git a/backend/internal/usersignup/service.go b/backend/internal/usersignup/service.go index 8b3307e0..26b23157 100644 --- a/backend/internal/usersignup/service.go +++ b/backend/internal/usersignup/service.go @@ -3,7 +3,6 @@ package usersignup import ( "context" "errors" - "fmt" "strings" "time" @@ -38,12 +37,7 @@ func newService(deps Dependencies) *Service { } } -func (s *Service) SignUp(ctx context.Context, signupData signUpDto, ipAddress, userAgent string) (model.User, string, error) { - config, err := appconfig.FromCtx(ctx) - if err != nil { - return model.User{}, "", fmt.Errorf("error loading app configuration: %w", err) - } - +func (s *Service) SignUp(ctx context.Context, config *appconfig.AppConfigModel, signupData signUpDto, ipAddress, userAgent string) (model.User, string, error) { tx := s.db.Begin() defer func() { tx.Rollback() @@ -91,7 +85,7 @@ func (s *Service) SignUp(ctx context.Context, signupData signUpDto, ipAddress, u EmailVerified: config.EmailsVerified.IsTrue(), } - user, err := s.userCreator.CreateUserInternal(ctx, userToCreate, false, tx) + user, err := s.userCreator.CreateUserInternal(ctx, config, userToCreate, false, tx) if err != nil { return model.User{}, "", err } @@ -126,12 +120,7 @@ func (s *Service) SignUp(ctx context.Context, signupData signUpDto, ipAddress, u return user, accessToken, nil } -func (s *Service) SignUpInitialAdmin(ctx context.Context, signUpData signUpDto) (model.User, string, error) { - config, err := appconfig.FromCtx(ctx) - if err != nil { - return model.User{}, "", fmt.Errorf("error loading app configuration: %w", err) - } - +func (s *Service) SignUpInitialAdmin(ctx context.Context, config *appconfig.AppConfigModel, signUpData signUpDto) (model.User, string, error) { tx := s.db.Begin() defer func() { tx.Rollback() @@ -154,7 +143,7 @@ func (s *Service) SignUpInitialAdmin(ctx context.Context, signUpData signUpDto) IsAdmin: true, } - user, err := s.userCreator.CreateUserInternal(ctx, userToCreate, false, tx) + user, err := s.userCreator.CreateUserInternal(ctx, config, userToCreate, false, tx) if err != nil { return model.User{}, "", err } diff --git a/backend/internal/webauthn/handler.go b/backend/internal/webauthn/handler.go index 7d616a5f..77d9e9a9 100644 --- a/backend/internal/webauthn/handler.go +++ b/backend/internal/webauthn/handler.go @@ -7,23 +7,29 @@ import ( "github.com/gin-gonic/gin" "github.com/go-webauthn/webauthn/protocol" - "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/utils/cookie" ) type handler struct { - service *Service + service *Service + appConfig AppConfigResolver } -func newHandler(service *Service) *handler { - return &handler{service: service} +func newHandler(service *Service, appConfig AppConfigResolver) *handler { + return &handler{service: service, appConfig: appConfig} } func (h *handler) beginRegistration(c *gin.Context) { + dbConfig, err := h.appConfig.GetConfig(c.Request.Context()) + if err != nil { + _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) + return + } + userID := c.GetString("userID") - options, err := h.service.BeginRegistration(c.Request.Context(), userID) + options, err := h.service.BeginRegistration(c.Request.Context(), dbConfig, userID) if err != nil { _ = c.Error(err) return @@ -68,7 +74,7 @@ func (h *handler) beginLogin(c *gin.Context) { } func (h *handler) verifyLogin(c *gin.Context) { - dbConfig, err := appconfig.FromCtx(c.Request.Context()) + dbConfig, err := h.appConfig.GetConfig(c.Request.Context()) if err != nil { _ = c.Error(fmt.Errorf("error loading app configuration: %w", err)) return @@ -86,7 +92,7 @@ func (h *handler) verifyLogin(c *gin.Context) { return } - user, token, err := h.service.VerifyLogin(c.Request.Context(), sessionID, credentialAssertionData, c.ClientIP(), c.Request.UserAgent()) + user, token, err := h.service.VerifyLogin(c.Request.Context(), dbConfig, sessionID, credentialAssertionData, c.ClientIP(), c.Request.UserAgent()) if err != nil { _ = c.Error(err) return diff --git a/backend/internal/webauthn/module.go b/backend/internal/webauthn/module.go index 0da68c53..5271bf69 100644 --- a/backend/internal/webauthn/module.go +++ b/backend/internal/webauthn/module.go @@ -23,12 +23,18 @@ type AuditLogger interface { CreateNewSignInWithEmail(ctx context.Context, ipAddress, userAgent, userID string, tx *gorm.DB, dbConfig *appconfig.AppConfigModel) model.AuditLog } +// AppConfigResolver loads the current application configuration, so handlers can pass it explicitly to the service methods that need it +type AppConfigResolver interface { + GetConfig(ctx context.Context) (*appconfig.AppConfigModel, error) +} + type Dependencies struct { DB *gorm.DB AppURL string - Signer TokenService - AuditLog AuditLogger + Signer TokenService + AuditLog AuditLogger + AppConfig AppConfigResolver } type Module struct { @@ -44,7 +50,7 @@ func New(deps Dependencies) (*Module, error) { return &Module{ service: service, - handler: newHandler(service), + handler: newHandler(service, deps.AppConfig), }, nil } diff --git a/backend/internal/webauthn/service.go b/backend/internal/webauthn/service.go index 00296c6f..8e387ca6 100644 --- a/backend/internal/webauthn/service.go +++ b/backend/internal/webauthn/service.go @@ -67,11 +67,8 @@ func newService(deps Dependencies) (*Service, error) { }, nil } -func (s *Service) BeginRegistration(ctx context.Context, userID string) (*PublicKeyCredentialCreationOptions, error) { - err := s.updateWebAuthnConfig(ctx) - if err != nil { - return nil, err - } +func (s *Service) BeginRegistration(ctx context.Context, dbConfig *appconfig.AppConfigModel, userID string) (*PublicKeyCredentialCreationOptions, error) { + s.updateWebAuthnConfig(dbConfig) tx := s.db.Begin() defer func() { @@ -79,7 +76,7 @@ func (s *Service) BeginRegistration(ctx context.Context, userID string) (*Public }() var user model.User - err = tx. + err := tx. WithContext(ctx). Preload("Credentials"). Find(&user, "id = ?", userID). @@ -232,12 +229,7 @@ func (s *Service) BeginLogin(ctx context.Context) (*PublicKeyCredentialRequestOp }, nil } -func (s *Service) VerifyLogin(ctx context.Context, sessionID string, credentialAssertionData *protocol.ParsedCredentialAssertionData, ipAddress, userAgent string) (model.User, string, error) { - dbConfig, err := appconfig.FromCtx(ctx) - if err != nil { - return model.User{}, "", fmt.Errorf("error loading app configuration: %w", err) - } - +func (s *Service) VerifyLogin(ctx context.Context, dbConfig *appconfig.AppConfigModel, sessionID string, credentialAssertionData *protocol.ParsedCredentialAssertionData, ipAddress, userAgent string) (model.User, string, error) { tx := s.db.Begin() defer func() { tx.Rollback() @@ -245,7 +237,7 @@ func (s *Service) VerifyLogin(ctx context.Context, sessionID string, credentialA // Load & delete the session row var storedSession WebauthnSession - err = tx. + err := tx. WithContext(ctx). Clauses(clause.Returning{}). Delete(&storedSession, "id = ?", sessionID). @@ -383,14 +375,8 @@ func (s *Service) UpdateCredential(ctx context.Context, userID, credentialID, na } // updateWebAuthnConfig updates the WebAuthn configuration with the app name as it can change during runtime -func (s *Service) updateWebAuthnConfig(ctx context.Context) error { - dbConfig, err := appconfig.FromCtx(ctx) - if err != nil { - return fmt.Errorf("error loading app configuration: %w", err) - } - +func (s *Service) updateWebAuthnConfig(dbConfig *appconfig.AppConfigModel) { s.webAuthn.Config.RPDisplayName = dbConfig.AppName.String() - return nil } func (s *Service) CreateReauthenticationTokenWithAccessToken(ctx context.Context, accessToken string) (string, error) { diff --git a/backend/internal/webauthn/service_test.go b/backend/internal/webauthn/service_test.go index ef4a9ec2..1ccfeece 100644 --- a/backend/internal/webauthn/service_test.go +++ b/backend/internal/webauthn/service_test.go @@ -128,9 +128,7 @@ func TestWebAuthnDisplayNameUsesRequestConfig(t *testing.T) { require.NoError(t, err) require.Equal(t, defaultRPDisplayName, service.webAuthn.Config.RPDisplayName) - ctx := appconfig.NewTestContext(t.Context(), &appconfig.AppConfigModel{AppName: "Custom App"}) - err = service.updateWebAuthnConfig(ctx) - require.NoError(t, err) + service.updateWebAuthnConfig(&appconfig.AppConfigModel{AppName: "Custom App"}) require.Equal(t, "Custom App", service.webAuthn.Config.RPDisplayName) }