From 0e0187dbc6fec670996da929a8ee573f61151c25 Mon Sep 17 00:00:00 2001 From: Claude Date: Sun, 19 Jul 2026 21:33:10 +0000 Subject: [PATCH] refactor: pass app config as explicit argument instead of request context Remove the AppConfigMiddleware that stored an app-config resolver in the request context and the FromCtx helper that read it back downstream. Handlers now load the app config from AppConfigService and pass it as an explicit argument to the service methods that need it, which then forward it further down the call chain. The webauthn and usersignup modules gain an AppConfigResolver dependency so their handlers can load the config the same way. The email SendEmail and LDAP SyncAll helpers drop their context-based variants in favor of the explicit-config versions. Co-Authored-By: Claude Opus 4.8 Claude-Session: https://claude.ai/code/session_01YZ6SoJpmnLggZqactXxrak --- .../appconfig/app_config_middleware.go | 64 ------------------- backend/internal/appconfig/testing_unit.go | 14 +--- .../internal/bootstrap/router_bootstrap.go | 10 ++- .../internal/bootstrap/services_bootstrap.go | 10 +-- .../controller/app_config_controller.go | 20 ++++-- .../internal/controller/user_controller.go | 56 +++++++++++++--- .../controller/user_group_controller.go | 22 ++++++- backend/internal/job/api_key_expiry_job.go | 2 +- backend/internal/job/ldap_job.go | 2 +- backend/internal/service/audit_log_service.go | 2 +- backend/internal/service/e2etest_service.go | 6 +- backend/internal/service/email_service.go | 17 ++--- backend/internal/service/ldap_service.go | 24 ++----- backend/internal/service/ldap_service_test.go | 8 +-- .../service/one_time_access_service.go | 33 +++------- .../service/one_time_access_service_test.go | 4 +- .../internal/service/user_group_service.go | 20 ++---- backend/internal/service/user_service.go | 33 +++------- backend/internal/service/user_service_test.go | 6 +- backend/internal/usersignup/handler.go | 16 ++--- backend/internal/usersignup/module.go | 11 +++- backend/internal/usersignup/service.go | 19 ++---- backend/internal/webauthn/handler.go | 20 ++++-- backend/internal/webauthn/module.go | 12 +++- backend/internal/webauthn/service.go | 26 ++------ backend/internal/webauthn/service_test.go | 4 +- 26 files changed, 195 insertions(+), 266 deletions(-) delete mode 100644 backend/internal/appconfig/app_config_middleware.go 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) }