diff --git a/backend/go.mod b/backend/go.mod index 9d5c0719..8a0e04cc 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -17,7 +17,6 @@ require ( github.com/fsnotify/fsnotify v1.10.1 github.com/gin-contrib/slog v1.2.1 github.com/gin-gonic/gin v1.12.0 - github.com/go-co-op/gocron/v2 v2.22.0 github.com/go-jose/go-jose/v4 v4.1.4 github.com/go-ldap/ldap/v3 v3.4.14 github.com/go-playground/validator/v10 v10.30.3 @@ -137,7 +136,6 @@ require ( github.com/jaegertracing/jaeger-idl v0.9.0 // indirect github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect - github.com/jonboulle/clockwork v0.5.0 // indirect github.com/json-iterator/go v1.1.12 // indirect github.com/klauspost/cpuid/v2 v2.4.0 // indirect github.com/leodido/go-urn v1.5.0 // indirect diff --git a/backend/go.sum b/backend/go.sum index 3187167c..f9ecd8f6 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -153,8 +153,6 @@ github.com/gin-gonic/gin v1.12.0 h1:b3YAbrZtnf8N//yjKeU2+MQsh2mY5htkZidOM7O0wG8= github.com/gin-gonic/gin v1.12.0/go.mod h1:VxccKfsSllpKshkBWgVgRniFFAzFb9csfngsqANjnLc= github.com/go-asn1-ber/asn1-ber v1.5.8 h1:H9AZkK22UOmfX8J84ubyaZxKJZ3FMHVwn8swoMML7iQ= github.com/go-asn1-ber/asn1-ber v1.5.8/go.mod h1:hEBeB/ic+5LoWskz+yKT7vGhhPYkProFKoKdwZRWMe0= -github.com/go-co-op/gocron/v2 v2.22.0 h1:uEuH2F7k7VoESb1BYSaffuuV+T0kkpzsC0aXk7/z79I= -github.com/go-co-op/gocron/v2 v2.22.0/go.mod h1:hiH/U9RMhTi1BBZJmef9s3KC9QwhpBF6PFrvUKaXY9M= github.com/go-errors/errors v1.0.1/go.mod h1:f4zRHt4oKfwPJE5k8C9vpYG+aDHdBFUsgrm6/TyX73Q= github.com/go-errors/errors v1.0.2/go.mod h1:psDX2osz5VnTOnFWbDeWwS7yejl+uV3FEWEp4lssFEs= github.com/go-errors/errors v1.1.1/go.mod h1:psDX2osz5VnTOnFWbDeWwS7yejl+uV3FEWEp4lssFEs= @@ -282,8 +280,6 @@ github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0= github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4= -github.com/jonboulle/clockwork v0.5.0 h1:Hyh9A8u51kptdkR+cqRpT1EebBwTn1oK9YfGYbdFz6I= -github.com/jonboulle/clockwork v0.5.0/go.mod h1:3mZlmanh0g2NDKO5TWZVJAfofYk64M7XN3SzBPjZF60= github.com/json-iterator/go v1.1.12 h1:PV8peI4a0ysnczrg+LtxykD8LfKY9ML6u2jnxaEnrnM= github.com/json-iterator/go v1.1.12/go.mod h1:e30LSqwooZae/UwlEbR2852Gd8hjQvJoHmT4TnhNGBo= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= diff --git a/backend/internal/apikey/expiry_job_test.go b/backend/internal/apikey/expiry_job_test.go index d8fd4124..281fa59f 100644 --- a/backend/internal/apikey/expiry_job_test.go +++ b/backend/internal/apikey/expiry_job_test.go @@ -50,9 +50,10 @@ func TestModuleRegistersAPIKeyExpiryCronJob(t *testing.T) { func TestAPIKeyExpiryCronJobNotifiesAndMarksExpiringKeys(t *testing.T) { db := testutils.NewDatabaseForTest(t) + userEmail := "expiry-job@example.com" user := model.User{ Username: "expiry-job-user", - Email: new("expiry-job@example.com"), + Email: &userEmail, FirstName: "Expiry", LastName: "Job", DisplayName: "Expiry Job", diff --git a/backend/internal/bootstrap/bootstrap.go b/backend/internal/bootstrap/bootstrap.go index c3635f13..1aaa7984 100644 --- a/backend/internal/bootstrap/bootstrap.go +++ b/backend/internal/bootstrap/bootstrap.go @@ -16,7 +16,6 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/instanceid" - "github.com/pocket-id/pocket-id/backend/internal/job" "github.com/pocket-id/pocket-id/backend/internal/storage" ) @@ -75,12 +74,6 @@ func Bootstrap(ctx context.Context) error { return fmt.Errorf("failed to initialize application images: %w", err) } - // Init the scheduler - scheduler, err := job.NewScheduler() - if err != nil { - return fmt.Errorf("failed to create job scheduler: %w", err) - } - // Init the actors // The actor host is created and started before the services, so services can depend on it once it's ready actorsOpts := NewActorsOpts{ @@ -102,7 +95,7 @@ func Bootstrap(ctx context.Context) error { services = append(services, actorsRun) // Create all services - svc, err := initServices(ctx, db, instanceID, actors, httpClient, imageExtensions, fileStorage, scheduler) + svc, err := initServices(ctx, db, instanceID, actors, httpClient, imageExtensions, fileStorage) if err != nil { return fmt.Errorf("failed to initialize services: %w", err) } @@ -110,18 +103,10 @@ func Bootstrap(ctx context.Context) error { // Migrate the pre-actor signup tokens into their actors, once the actor host is ready services = append(services, actorsReady.Await(svc.userSignUpModule.RunSignupTokenMigration)) - // Register scheduled jobs, only in non-test mode + // These services are only registered in non-test mode if common.EnvConfig.AppEnv != "test" { - err = registerScheduledJobs(ctx, svc, scheduler) - if err != nil { - return fmt.Errorf("failed to register scheduled jobs: %w", err) - } - // Refresh the GeoLite database (this is cached per each replica) services = append(services, svc.geoLiteModule.Run) - - // The scheduler must wait on the actor host being ready, since jobs invoke actors - services = append(services, actorsReady.Await(scheduler.Run)) } // Init the router diff --git a/backend/internal/bootstrap/router_bootstrap.go b/backend/internal/bootstrap/router_bootstrap.go index bd9ee767..07c5d354 100644 --- a/backend/internal/bootstrap/router_bootstrap.go +++ b/backend/internal/bootstrap/router_bootstrap.go @@ -177,7 +177,7 @@ func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, rateLimitServices svc.apiModule.RegisterRoutes(apiGroup, authMiddleware.Add()) controller.NewCustomClaimController(apiGroup, authMiddleware, svc.customClaimService) controller.NewVersionController(apiGroup, authMiddleware, svc.versionService) - controller.NewScimController(apiGroup, authMiddleware, svc.scimService) + svc.scimSyncModule.RegisterRoutes(apiGroup, authMiddleware.Add()) svc.userSignUpModule.RegisterRoutes(apiGroup, authMiddleware.Add(), rateLimitMiddleware.Add(middleware.RateLimitSignup), diff --git a/backend/internal/bootstrap/scheduler_bootstrap.go b/backend/internal/bootstrap/scheduler_bootstrap.go deleted file mode 100644 index 2abf2136..00000000 --- a/backend/internal/bootstrap/scheduler_bootstrap.go +++ /dev/null @@ -1,17 +0,0 @@ -package bootstrap - -import ( - "context" - "fmt" - - "github.com/pocket-id/pocket-id/backend/internal/job" -) - -func registerScheduledJobs(ctx context.Context, svc *services, scheduler *job.Scheduler) error { - err := scheduler.RegisterScimJobs(ctx, svc.scimService) - if err != nil { - return fmt.Errorf("failed to register SCIM scheduler job: %w", err) - } - - return nil -} diff --git a/backend/internal/bootstrap/services_bootstrap.go b/backend/internal/bootstrap/services_bootstrap.go index 7c619c2d..5c653c47 100644 --- a/backend/internal/bootstrap/services_bootstrap.go +++ b/backend/internal/bootstrap/services_bootstrap.go @@ -15,10 +15,10 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/email" "github.com/pocket-id/pocket-id/backend/internal/emailverification" "github.com/pocket-id/pocket-id/backend/internal/geolite" - "github.com/pocket-id/pocket-id/backend/internal/job" "github.com/pocket-id/pocket-id/backend/internal/ldapsync" "github.com/pocket-id/pocket-id/backend/internal/oidc" "github.com/pocket-id/pocket-id/backend/internal/onetimeaccess" + "github.com/pocket-id/pocket-id/backend/internal/scimsync" "github.com/pocket-id/pocket-id/backend/internal/service" "github.com/pocket-id/pocket-id/backend/internal/storage" "github.com/pocket-id/pocket-id/backend/internal/usersignup" @@ -33,7 +33,6 @@ type services struct { geoLiteModule *geolite.Module auditLogService *service.AuditLogService jwtService *service.JwtService - scimService *service.ScimService userService *service.UserService customClaimService *service.CustomClaimService oidcService *service.OidcService @@ -45,6 +44,7 @@ type services struct { auditLogsModule *auditlogs.Module deviceLoginModule *devicelogin.Module ldapSyncModule *ldapsync.Module + scimSyncModule *scimsync.Module oidcModule *oidc.Module webauthnModule *webauthn.Module userSignUpModule *usersignup.Module @@ -63,7 +63,6 @@ func initServices( httpClient *http.Client, imageExtensions map[string]string, fileStorage storage.FileStorage, - scheduler *job.Scheduler, ) (svc *services, err error) { svc = &services{ actors: actors, @@ -138,7 +137,16 @@ func initServices( return nil, fmt.Errorf("failed to create device login module: %w", err) } - svc.scimService = service.NewScimService(db, scheduler, httpClient) + svc.scimSyncModule, err = scimsync.New(scimsync.Dependencies{ + DB: db, + Actors: actors, + HTTPClient: httpClient, + // Disable in test environment + ScheduleDisabled: common.EnvConfig.AppEnv.IsTest(), + }) + if err != nil { + return nil, fmt.Errorf("failed to create SCIM sync module: %w", err) + } svc.apiModule = api.New(api.Dependencies{DB: db, Issuer: common.EnvConfig.AppURL}) @@ -165,13 +173,13 @@ func initServices( return nil, fmt.Errorf("failed to create OIDC module: %w", err) } - svc.oidcService, err = service.NewOidcService(db, svc.jwtService, svc.oidcModule.Preview, svc.oidcModule, svc.scimService, httpClient, fileStorage) + svc.oidcService, err = service.NewOidcService(db, svc.jwtService, svc.oidcModule.Preview, svc.oidcModule, svc.scimSyncModule, httpClient, fileStorage) if err != nil { return nil, fmt.Errorf("failed to create OIDC service: %w", err) } - svc.userGroupService = service.NewUserGroupService(db, svc.scimService) - svc.userService = service.NewUserService(db, svc.jwtService, svc.auditLogService, svc.customClaimService, svc.appImagesService, svc.scimService, fileStorage) + svc.userGroupService = service.NewUserGroupService(db, svc.scimSyncModule) + svc.userService = service.NewUserService(db, svc.jwtService, svc.auditLogService, svc.customClaimService, svc.appImagesService, svc.scimSyncModule, fileStorage) svc.ldapSyncModule, err = ldapsync.New(ldapsync.Dependencies{ DB: db, @@ -181,6 +189,7 @@ func initServices( Users: svc.userService, Groups: svc.userGroupService, AppConfig: svc.appConfigService, + ScimSync: svc.scimSyncModule, // Disable in test environment ScheduleDisabled: common.EnvConfig.AppEnv.IsTest(), }) @@ -207,6 +216,7 @@ func initServices( AuditLog: svc.auditLogService, UserCreator: svc.userService, AppConfig: svc.appConfigService, + ScimSync: svc.scimSyncModule, }) if err != nil { return nil, fmt.Errorf("failed to create user signup module: %w", err) diff --git a/backend/internal/cmds/encryption_key_rotate.go b/backend/internal/cmds/encryption_key_rotate.go index e620ac3c..dddd0a7a 100644 --- a/backend/internal/cmds/encryption_key_rotate.go +++ b/backend/internal/cmds/encryption_key_rotate.go @@ -12,8 +12,8 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/bootstrap" "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/instanceid" - "github.com/pocket-id/pocket-id/backend/internal/model" datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" + "github.com/pocket-id/pocket-id/backend/internal/scimsync" "github.com/pocket-id/pocket-id/backend/internal/utils" jwkutils "github.com/pocket-id/pocket-id/backend/internal/utils/jwk" ) @@ -156,7 +156,9 @@ type scimTokenRow struct { func rotateScimTokens(db *gorm.DB, oldEncKey []byte, newEncKey []byte) error { var rows []scimTokenRow - err := db.Model(&model.ScimServiceProvider{}).Select("id, token").Scan(&rows).Error + err := db.Model(&scimsync.ServiceProvider{}). + Select("id, token"). + Scan(&rows).Error if err != nil { return fmt.Errorf("failed to list SCIM service providers: %w", err) } @@ -176,7 +178,9 @@ func rotateScimTokens(db *gorm.DB, oldEncKey []byte, newEncKey []byte) error { return fmt.Errorf("failed to encrypt SCIM token for provider %s: %w", row.ID, err) } - err = db.Model(&model.ScimServiceProvider{}).Where("id = ?", row.ID).Update("token", encValue).Error + err = db.Model(&scimsync.ServiceProvider{}). + Where("id = ?", row.ID). + Update("token", encValue).Error if err != nil { return fmt.Errorf("failed to update SCIM token for provider %s: %w", row.ID, err) } diff --git a/backend/internal/cmds/encryption_key_rotate_test.go b/backend/internal/cmds/encryption_key_rotate_test.go index a22f0e46..e551943b 100644 --- a/backend/internal/cmds/encryption_key_rotate_test.go +++ b/backend/internal/cmds/encryption_key_rotate_test.go @@ -9,8 +9,8 @@ import ( "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/instanceid" - "github.com/pocket-id/pocket-id/backend/internal/model" datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" + "github.com/pocket-id/pocket-id/backend/internal/scimsync" jwkutils "github.com/pocket-id/pocket-id/backend/internal/utils/jwk" testingutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing" ) @@ -77,7 +77,10 @@ func TestEncryptionKeyRotate(t *testing.T) { require.NotNil(t, rotatedKey) var storedToken string - err = db.Model(&model.ScimServiceProvider{}).Where("id = ?", "scim-1").Pluck("token", &storedToken).Error + err = db.Model(&scimsync.ServiceProvider{}). + Where("id = ?", "scim-1"). + Pluck("token", &storedToken). + Error require.NoError(t, err) newEncKey, err := datatype.DeriveEncryptedStringKey(newKey) diff --git a/backend/internal/controller/oidc_controller.go b/backend/internal/controller/oidc_controller.go index 9df00fac..d54d9a86 100644 --- a/backend/internal/controller/oidc_controller.go +++ b/backend/internal/controller/oidc_controller.go @@ -51,8 +51,6 @@ func NewOidcController(group *gin.RouterGroup, authMiddleware *middleware.AuthMi group.GET("/oidc/users/me/clients", authMiddleware.WithAdminNotRequired().Add(), httpserver.Handle(oc.listOwnAccessibleClientsHandler)) - group.GET("/oidc/clients/:id/scim-service-provider", authMiddleware.Add(), httpserver.Handle(oc.getClientScimServiceProviderHandler)) - } type OidcController struct { @@ -607,29 +605,3 @@ func (oc *OidcController) getClientPreviewHandler(c *gin.Context) error { c.JSON(http.StatusOK, preview) return nil } - -// getClientScimServiceProviderHandler godoc -// @Summary Get SCIM service provider -// @Description Get the SCIM service provider configuration for an OIDC client -// @Tags OIDC -// @Produce json -// @Param id path string true "Client ID" -// @Success 200 {object} dto.ScimServiceProviderDTO "SCIM service provider configuration" -// @Failure default {object} dto.ErrorDto "Error" -// @Router /api/oidc/clients/{id}/scim-service-provider [get] -func (oc *OidcController) getClientScimServiceProviderHandler(c *gin.Context) error { - clientID := c.Param("id") - - provider, err := oc.oidcService.GetClientScimServiceProvider(c.Request.Context(), clientID) - if err != nil { - return err - } - - var providerDto dto.ScimServiceProviderDTO - if err := dto.MapStruct(provider, &providerDto); err != nil { - return err - } - - c.JSON(http.StatusOK, providerDto) - return nil -} diff --git a/backend/internal/controller/scim_controller.go b/backend/internal/controller/scim_controller.go deleted file mode 100644 index 14ab9828..00000000 --- a/backend/internal/controller/scim_controller.go +++ /dev/null @@ -1,123 +0,0 @@ -package controller - -import ( - "net/http" - - "github.com/gin-gonic/gin" - "github.com/pocket-id/pocket-id/backend/internal/dto" - "github.com/pocket-id/pocket-id/backend/internal/httpserver" - "github.com/pocket-id/pocket-id/backend/internal/middleware" - "github.com/pocket-id/pocket-id/backend/internal/service" -) - -func NewScimController(group *gin.RouterGroup, authMiddleware *middleware.AuthMiddleware, scimService *service.ScimService) { - ugc := ScimController{ - scimService: scimService, - } - - group.POST("/scim/service-provider", authMiddleware.Add(), httpserver.Handle(ugc.createServiceProviderHandler)) - group.POST("/scim/service-provider/:id/sync", authMiddleware.Add(), httpserver.Handle(ugc.syncServiceProviderHandler)) - group.PUT("/scim/service-provider/:id", authMiddleware.Add(), httpserver.Handle(ugc.updateServiceProviderHandler)) - group.DELETE("/scim/service-provider/:id", authMiddleware.Add(), httpserver.Handle(ugc.deleteServiceProviderHandler)) -} - -type ScimController struct { - scimService *service.ScimService -} - -// syncServiceProviderHandler godoc -// @Summary Sync SCIM service provider -// @Description Trigger synchronization for a SCIM service provider -// @Tags SCIM -// @Param id path string true "Service Provider ID" -// @Success 200 "OK" -// @Failure default {object} dto.ErrorDto "Error" -// @Router /api/scim/service-provider/{id}/sync [post] -func (c *ScimController) syncServiceProviderHandler(ctx *gin.Context) error { - err := c.scimService.SyncServiceProvider(ctx.Request.Context(), ctx.Param("id")) - if err != nil { - return err - } - - ctx.Status(http.StatusOK) - return nil -} - -// createServiceProviderHandler godoc -// @Summary Create SCIM service provider -// @Description Create a new SCIM service provider -// @Tags SCIM -// @Accept json -// @Produce json -// @Param serviceProvider body dto.ScimServiceProviderCreateDTO true "SCIM service provider information" -// @Success 201 {object} dto.ScimServiceProviderDTO "Created SCIM service provider" -// @Failure default {object} dto.ErrorDto "Error" -// @Router /api/scim/service-provider [post] -func (c *ScimController) createServiceProviderHandler(ctx *gin.Context) error { - var input dto.ScimServiceProviderCreateDTO - if err := httpserver.BindJSON(ctx, &input); err != nil { - return err - } - - provider, err := c.scimService.CreateServiceProvider(ctx.Request.Context(), &input) - if err != nil { - return err - } - - var providerDTO dto.ScimServiceProviderDTO - if err := dto.MapStruct(provider, &providerDTO); err != nil { - return err - } - - ctx.JSON(http.StatusCreated, providerDTO) - return nil -} - -// updateServiceProviderHandler godoc -// @Summary Update SCIM service provider -// @Description Update an existing SCIM service provider -// @Tags SCIM -// @Accept json -// @Produce json -// @Param id path string true "Service Provider ID" -// @Param serviceProvider body dto.ScimServiceProviderCreateDTO true "SCIM service provider information" -// @Success 200 {object} dto.ScimServiceProviderDTO "Updated SCIM service provider" -// @Failure default {object} dto.ErrorDto "Error" -// @Router /api/scim/service-provider/{id} [put] -func (c *ScimController) updateServiceProviderHandler(ctx *gin.Context) error { - var input dto.ScimServiceProviderCreateDTO - if err := httpserver.BindJSON(ctx, &input); err != nil { - return err - } - - provider, err := c.scimService.UpdateServiceProvider(ctx.Request.Context(), ctx.Param("id"), &input) - if err != nil { - return err - } - - var providerDTO dto.ScimServiceProviderDTO - if err := dto.MapStruct(provider, &providerDTO); err != nil { - return err - } - - ctx.JSON(http.StatusOK, providerDTO) - return nil -} - -// deleteServiceProviderHandler godoc -// @Summary Delete SCIM service provider -// @Description Delete a SCIM service provider by ID -// @Tags SCIM -// @Param id path string true "Service Provider ID" -// @Success 204 "No Content" -// @Failure default {object} dto.ErrorDto "Error" -// @Router /api/scim/service-provider/{id} [delete] -func (c *ScimController) deleteServiceProviderHandler(ctx *gin.Context) error { - err := c.scimService.DeleteServiceProvider(ctx.Request.Context(), ctx.Param("id")) - if err != nil { - return err - } - - ctx.Status(http.StatusNoContent) - return nil -} diff --git a/backend/internal/job/scheduler.go b/backend/internal/job/scheduler.go deleted file mode 100644 index ad7f007e..00000000 --- a/backend/internal/job/scheduler.go +++ /dev/null @@ -1,165 +0,0 @@ -package job - -import ( - "context" - "errors" - "fmt" - "log/slog" - "time" - - backoff "github.com/cenkalti/backoff/v5" - "github.com/go-co-op/gocron/v2" - "github.com/google/uuid" - "go.opentelemetry.io/otel/trace" - - "github.com/pocket-id/pocket-id/backend/internal/service" - "github.com/pocket-id/pocket-id/backend/internal/tracing" -) - -type Scheduler struct { - scheduler gocron.Scheduler -} - -func NewScheduler() (*Scheduler, error) { - scheduler, err := gocron.NewScheduler() - if err != nil { - return nil, fmt.Errorf("failed to create a new scheduler: %w", err) - } - - return &Scheduler{ - scheduler: scheduler, - }, nil -} - -func (s *Scheduler) RemoveJob(name string) error { - jobs := s.scheduler.Jobs() - - var errs []error - for _, job := range jobs { - if job.Name() == name { - err := s.scheduler.RemoveJob(job.ID()) - if err != nil { - errs = append(errs, fmt.Errorf("failed to dequeue job %q with ID %q: %w", name, job.ID().String(), err)) - } - } - } - - return errors.Join(errs...) -} - -// Run the scheduler. -// This function blocks until the context is canceled. -func (s *Scheduler) Run(ctx context.Context) error { - slog.Info("Starting job scheduler") - s.scheduler.Start() - - // Block until context is canceled - <-ctx.Done() - - err := s.scheduler.Shutdown() - if err != nil { - slog.Error("Error shutting down job scheduler", slog.Any("error", err)) - } else { - slog.Info("Job scheduler shut down") - } - - return nil -} - -func (s *Scheduler) RegisterJob(ctx context.Context, name string, def gocron.JobDefinition, jobFn func(ctx context.Context) error, opts service.RegisterJobOpts) error { - // Wrap the job in a handler that adds tracing and logging - jobFn = jobWithObservability(name, jobFn) - - // If a BackOff strategy is provided, wrap the job with retry logic - if opts.BackOff != nil { - jobFn = jobWithBackOff(jobFn, opts.BackOff) - } - - jobOptions := []gocron.JobOption{ - gocron.WithContext(ctx), - gocron.WithName(name), - } - - if opts.RunImmediately { - jobOptions = append(jobOptions, gocron.JobOption(gocron.WithStartImmediately())) - } - - jobOptions = append(jobOptions, opts.ExtraOptions...) - - _, err := s.scheduler.NewJob(def, gocron.NewTask(jobFn), jobOptions...) - if err != nil { - return fmt.Errorf("failed to register job %q: %w", name, err) - } - - return nil -} - -type ( - jobNameKey struct{} - jobIDKey struct{} - jobFn = func(ctx context.Context) error -) - -func jobWithObservability(jobName string, job jobFn) jobFn { - return func(ctx context.Context) error { - // Generate a random job ID - jobID := uuid.NewString() - - // Save in the context - ctx = context.WithValue(ctx, jobNameKey{}, jobName) - ctx = context.WithValue(ctx, jobIDKey{}, jobID) - - // Create a new context with the span - var err error - ctx, span := tracing.Start(ctx, "pocketid.job."+jobName, - trace.WithSpanKind(trace.SpanKindInternal), - trace.WithAttributes( - tracing.JobID(jobID), - ), - ) - defer tracing.End(span, err) - - // Log the start - logger := slog.With( - slog.String("name", jobName), - slog.String("jobID", jobID), - ) - start := time.Now() - logger.InfoContext(ctx, "Starting job") - - // Run the job - err = job(ctx) - d := time.Since(start) - if err != nil { - logger.ErrorContext(ctx, "Job failed", slog.Any("error", err), slog.Duration("duration", d)) - return err - } - - logger.InfoContext(ctx, "Job run successfully", slog.Duration("duration", d)) - return nil - } -} - -func jobWithBackOff(job jobFn, bo backoff.BackOff) jobFn { - return func(ctx context.Context) error { - jobName, _ := (ctx.Value(jobNameKey{})).(string) - jobID, _ := (ctx.Value(jobIDKey{})).(string) - - _, err := backoff.Retry( - ctx, - func() (struct{}, error) { - return struct{}{}, job(ctx) - }, - backoff.WithBackOff(bo), - backoff.WithNotify(func(err error, d time.Duration) { - slog.WarnContext(ctx, "Job failed, retrying", - slog.String("name", jobName), - slog.String("jobID", jobID), - slog.Any("error", err), - slog.Duration("retryIn", d), - ) - }), - ) - return err - } -} diff --git a/backend/internal/job/scim_job.go b/backend/internal/job/scim_job.go deleted file mode 100644 index 5c4336f6..00000000 --- a/backend/internal/job/scim_job.go +++ /dev/null @@ -1,25 +0,0 @@ -package job - -import ( - "context" - "time" - - "github.com/go-co-op/gocron/v2" - - "github.com/pocket-id/pocket-id/backend/internal/service" -) - -type ScimJobs struct { - scimService *service.ScimService -} - -func (s *Scheduler) RegisterScimJobs(ctx context.Context, scimService *service.ScimService) error { - jobs := &ScimJobs{scimService: scimService} - - // Register the job to run every hour (with some jitter) - return s.RegisterJob(ctx, "SyncScim", gocron.DurationJob(time.Hour), jobs.SyncScim, service.RegisterJobOpts{RunImmediately: true}) -} - -func (j *ScimJobs) SyncScim(ctx context.Context) error { - return j.scimService.SyncAll(ctx) -} diff --git a/backend/internal/ldapsync/module.go b/backend/internal/ldapsync/module.go index 81aafb66..ccfca920 100644 --- a/backend/internal/ldapsync/module.go +++ b/backend/internal/ldapsync/module.go @@ -36,6 +36,11 @@ type GroupSyncer interface { UpdateUsersInternal(ctx context.Context, id string, userIDs []string, tx *gorm.DB) (model.UserGroup, error) } +// ScimSyncScheduler schedules SCIM after the LDAP transaction has committed +type ScimSyncScheduler interface { + ScheduleSync(ctx context.Context) +} + type Dependencies struct { DB *gorm.DB Actors *local.Host @@ -45,6 +50,7 @@ type Dependencies struct { Users UserSyncer Groups GroupSyncer AppConfig appconfig.AppConfigResolver + ScimSync ScimSyncScheduler // ScheduleDisabled keeps the recurring sync from being armed // It's set in the test environment, where syncs are driven explicitly by the end-to-end tests diff --git a/backend/internal/ldapsync/service.go b/backend/internal/ldapsync/service.go index 511c087a..7e661cb1 100644 --- a/backend/internal/ldapsync/service.go +++ b/backend/internal/ldapsync/service.go @@ -35,6 +35,7 @@ type Service struct { httpClient *http.Client users UserSyncer groups GroupSyncer + scimSync ScimSyncScheduler fileStorage storage.FileStorage clientFactory func(dbConfig *appconfig.AppConfigModel) (ldapClient, error) } @@ -76,6 +77,7 @@ func newService(deps Dependencies) *Service { httpClient: deps.HTTPClient, users: deps.Users, groups: deps.Groups, + scimSync: deps.ScimSync, fileStorage: deps.FileStorage, } @@ -144,6 +146,11 @@ func (s *Service) SyncAll(ctx context.Context, dbConfig *appconfig.AppConfigMode return fmt.Errorf("failed to commit changes to database: %w", err) } + // Schedule downstream SCIM reconciliation only after the LDAP transaction releases its database locks + if s.scimSync != nil { + s.scimSync.ScheduleSync(ctx) + } + // Now that we've committed the transaction, we can perform operations on the storage layer // First, save all new pictures for _, sp := range savePictures { diff --git a/backend/internal/model/scim.go b/backend/internal/model/scim.go deleted file mode 100644 index 1d8209c5..00000000 --- a/backend/internal/model/scim.go +++ /dev/null @@ -1,14 +0,0 @@ -package model - -import datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" - -type ScimServiceProvider struct { - Base - - Endpoint string `sortable:"true"` - Token datatype.EncryptedString - LastSyncedAt *datatype.DateTime `sortable:"true"` - - OidcClientID string - OidcClient OidcClient `gorm:"foreignKey:OidcClientID;references:ID;"` -} diff --git a/backend/internal/scimsync/actor.go b/backend/internal/scimsync/actor.go new file mode 100644 index 00000000..ba7490f2 --- /dev/null +++ b/backend/internal/scimsync/actor.go @@ -0,0 +1,148 @@ +package scimsync + +import ( + "context" + "errors" + "fmt" + "log/slog" + "time" + + "github.com/italypaleale/francis/actor" + + "github.com/pocket-id/pocket-id/backend/internal/common" +) + +// The ScimSync singleton actor decides when the recurring and debounced SCIM synchronizations run + +// ActorType is the actor type for the SCIM sync actor +const ActorType = "ScimSync" + +const ( + // alarmRecurringSync runs the cluster-wide hourly synchronization + alarmRecurringSync = "recurring-sync" + // alarmScheduledSync runs the cluster-wide synchronization requested after a local change + alarmScheduledSync = "scheduled-sync" + + // methodScheduleSync moves the debounced synchronization five minutes past the latest change + methodScheduleSync = "schedule-sync" + + // recurringSyncInterval is how often the full synchronization runs, as the ISO8601 duration the alarm repetition expects + recurringSyncInterval = "PT1H" + // initialSyncDelay gives the application a moment to finish starting before the first synchronization + initialSyncDelay = 5 * time.Second + // scheduledSyncDelay debounces changes so several related writes produce one synchronization + scheduledSyncDelay = 5 * time.Minute + // alarmTimeout bounds the alarm operations performed by the actor + alarmTimeout = 10 * time.Second +) + +type syncer interface { + SyncAll(ctx context.Context) error +} + +// syncActor is the cluster-wide singleton that triggers SCIM synchronization +type syncActor struct { + log *slog.Logger + syncer syncer + scheduleDisabled bool + client actor.Client[struct{}] +} + +// NewActor returns the factory that allocates the SCIM sync actor +func NewActor(service *Service, scheduleDisabled bool) actor.Factory { + return newActor(service, scheduleDisabled) +} + +func newActor(syncer syncer, scheduleDisabled bool) actor.Factory { + return func(actorID string, actorService *actor.Service) actor.Actor { + return &syncActor{ + log: slog.With( + slog.String("scope", "actor"), + slog.String("actorType", ActorType), + ), + syncer: syncer, + scheduleDisabled: scheduleDisabled, + client: actor.NewActorClient[struct{}](ActorType, actorID, actorService), + } + } +} + +// Bootstrap implements actor.ActorBootstrapper +// The host drives it on every startup, routed to the single owning host, so it must stay idempotent +func (a *syncActor) Bootstrap(parentCtx context.Context, _ actor.Envelope) error { + ctx, cancel := context.WithTimeout(parentCtx, alarmTimeout) + defer cancel() + + // The test environment drives synchronization explicitly and must not inherit alarms from a previous run + if a.scheduleDisabled { + for _, name := range []string{alarmRecurringSync, alarmScheduledSync} { + err := a.client.DeleteAlarm(ctx, name) + if err != nil && !errors.Is(err, actor.ErrAlarmNotFound) { + return fmt.Errorf("error deleting the SCIM sync alarm %q: %w", name, err) + } + } + + return nil + } + + // Replacing the alarm restores a missing schedule and applies interval changes after an upgrade + err := a.client.SetAlarm(ctx, alarmRecurringSync, actor.AlarmProperties{ + DueTime: time.Now().Add(initialSyncDelay), + Interval: recurringSyncInterval, + }) + if err != nil { + return fmt.Errorf("error setting the recurring SCIM sync alarm: %w", err) + } + + a.log.DebugContext(parentCtx, "Registered the recurring SCIM sync alarm", slog.String("interval", recurringSyncInterval)) + + return nil +} + +// Invoke implements actor.ActorInvoke +func (a *syncActor) Invoke(ctx context.Context, method string, _ actor.Envelope) (any, error) { + if method != methodScheduleSync { + return nil, common.ErrUnsupportedActorMethod{Method: method} + } + + if a.scheduleDisabled { + return nil, nil + } + + // Setting the same alarm replaces its due time, which debounces changes across every replica + err := a.client.SetAlarm(ctx, alarmScheduledSync, actor.AlarmProperties{ + DueTime: time.Now().Add(scheduledSyncDelay), + }) + if err != nil { + return nil, fmt.Errorf("error setting the scheduled SCIM sync alarm: %w", err) + } + + return nil, nil +} + +// Alarm implements actor.ActorAlarm +func (a *syncActor) Alarm(ctx context.Context, name string, _ actor.Envelope) error { + if name != alarmRecurringSync && name != alarmScheduledSync { + return fmt.Errorf("unsupported alarm '%s' for the %s actor", name, ActorType) + } + + a.sync(ctx) + + // A failed recurring sync must not surface as an error because exhausted alarm retries would remove the recurring schedule + // The sync method shows a log in case of failure + return nil +} + +// sync runs one full SCIM synchronization and records its outcome +func (a *syncActor) sync(ctx context.Context) { + a.log.InfoContext(ctx, "Starting the SCIM sync") + start := time.Now() + + err := a.syncer.SyncAll(ctx) + if err != nil { + a.log.ErrorContext(ctx, "SCIM sync failed, will try again on the next run", slog.Duration("duration", time.Since(start)), slog.Any("error", err)) + return + } + + a.log.InfoContext(ctx, "SCIM sync completed", slog.Duration("duration", time.Since(start))) +} diff --git a/backend/internal/scimsync/actor_test.go b/backend/internal/scimsync/actor_test.go new file mode 100644 index 00000000..f8f4fb93 --- /dev/null +++ b/backend/internal/scimsync/actor_test.go @@ -0,0 +1,180 @@ +package scimsync + +import ( + "context" + "errors" + "sync/atomic" + "testing" + "time" + + "github.com/italypaleale/francis/actor" + "github.com/italypaleale/francis/host/local" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing" +) + +type fakeSyncer struct { + calls atomic.Int32 + err error +} + +func (s *fakeSyncer) SyncAll(_ context.Context) error { + s.calls.Add(1) + return s.err +} + +func TestActorBootstrapArmsRecurringAlarm(t *testing.T) { + host, act := newSyncActorForTest(t, &fakeSyncer{}, false) + + require.NoError(t, act.Bootstrap(t.Context(), nil)) + + properties, err := host.GetAlarm(t.Context(), ActorType, actor.SingletonActorID, alarmRecurringSync) + require.NoError(t, err) + assert.Equal(t, recurringSyncInterval, properties.Interval) + assert.WithinDuration(t, time.Now().Add(initialSyncDelay), properties.DueTime, time.Second) +} + +func TestActorBootstrapIsIdempotent(t *testing.T) { + host, act := newSyncActorForTest(t, &fakeSyncer{}, false) + + require.NoError(t, act.Bootstrap(t.Context(), nil)) + require.NoError(t, act.Bootstrap(t.Context(), nil)) + + properties, err := host.GetAlarm(t.Context(), ActorType, actor.SingletonActorID, alarmRecurringSync) + require.NoError(t, err) + assert.Equal(t, recurringSyncInterval, properties.Interval) +} + +func TestActorBootstrapRemovesAutomaticAlarmsWhenDisabled(t *testing.T) { + host, act := newSyncActorForTest(t, &fakeSyncer{}, true) + + for _, name := range []string{alarmRecurringSync, alarmScheduledSync} { + require.NoError(t, host.SetAlarm(t.Context(), ActorType, actor.SingletonActorID, name, actor.AlarmProperties{ + DueTime: time.Now().Add(time.Hour), + })) + } + + require.NoError(t, act.Bootstrap(t.Context(), nil)) + + for _, name := range []string{alarmRecurringSync, alarmScheduledSync} { + _, err := host.GetAlarm(t.Context(), ActorType, actor.SingletonActorID, name) + require.ErrorIs(t, err, actor.ErrAlarmNotFound) + } +} + +func TestActorSchedulesDebouncedSync(t *testing.T) { + host, act := newSyncActorForTest(t, &fakeSyncer{}, false) + + _, err := act.Invoke(t.Context(), methodScheduleSync, nil) + require.NoError(t, err) + + first, err := host.GetAlarm(t.Context(), ActorType, actor.SingletonActorID, alarmScheduledSync) + require.NoError(t, err) + assert.WithinDuration(t, time.Now().Add(scheduledSyncDelay), first.DueTime, time.Second) + assert.Empty(t, first.Interval) + + _, err = act.Invoke(t.Context(), methodScheduleSync, nil) + require.NoError(t, err) + + second, err := host.GetAlarm(t.Context(), ActorType, actor.SingletonActorID, alarmScheduledSync) + require.NoError(t, err) + assert.False(t, second.DueTime.Before(first.DueTime)) +} + +func TestActorDoesNotScheduleDebouncedSyncWhenDisabled(t *testing.T) { + host, act := newSyncActorForTest(t, &fakeSyncer{}, true) + + _, err := act.Invoke(t.Context(), methodScheduleSync, nil) + require.NoError(t, err) + + _, err = host.GetAlarm(t.Context(), ActorType, actor.SingletonActorID, alarmScheduledSync) + require.ErrorIs(t, err, actor.ErrAlarmNotFound) +} + +func TestActorAlarmsRunSync(t *testing.T) { + syncer := &fakeSyncer{} + _, act := newSyncActorForTest(t, syncer, false) + + require.NoError(t, act.Alarm(t.Context(), alarmRecurringSync, nil)) + require.NoError(t, act.Alarm(t.Context(), alarmScheduledSync, nil)) + assert.EqualValues(t, 2, syncer.calls.Load()) +} + +func TestActorAlarmSwallowsSyncFailures(t *testing.T) { + syncer := &fakeSyncer{err: errors.New("provider unavailable")} + _, act := newSyncActorForTest(t, syncer, false) + + require.NoError(t, act.Alarm(t.Context(), alarmRecurringSync, nil)) + assert.EqualValues(t, 1, syncer.calls.Load()) +} + +func TestActorRejectsUnknownOperations(t *testing.T) { + _, act := newSyncActorForTest(t, &fakeSyncer{}, false) + + _, err := act.Invoke(t.Context(), "unknown", nil) + require.Error(t, err) + + err = act.Alarm(t.Context(), "unknown", nil) + require.Error(t, err) + assert.ErrorContains(t, err, "unsupported alarm") +} + +func TestRegisteredSingletonBootstrapsAndFires(t *testing.T) { + syncer := &fakeSyncer{} + host := testutils.NewActorHostForTest(t, + func(t *testing.T, host *local.Host) { + err := host.RegisterSingletonActor(ActorType, newActor(syncer, false)) + require.NoError(t, err) + }, + local.WithAlarmsPollInterval(5*time.Minute), + local.WithAlarmsFetchAheadInterval(5*time.Minute), + ) + + // The host bootstraps singleton actors asynchronously after it becomes ready + require.Eventually(t, func() bool { + _, err := host.GetAlarm(t.Context(), ActorType, actor.SingletonActorID, alarmRecurringSync) + return err == nil + }, 10*time.Second, 20*time.Millisecond) + + // The first occurrence is delivered without waiting for the normal five-minute alarm poll + require.Eventually(t, func() bool { + return syncer.calls.Load() > 0 + }, initialSyncDelay+30*time.Second, 50*time.Millisecond) +} + +func TestModuleRegistersSingletonAndSchedulesClusterWideSync(t *testing.T) { + var module *Module + host := testutils.NewActorHostForTest(t, func(t *testing.T, host *local.Host) { + var err error + module, err = New(Dependencies{ + DB: testutils.NewDatabaseForTest(t), + Actors: host, + }) + require.NoError(t, err) + }) + + // The host bootstraps singleton actors asynchronously after it becomes ready + require.Eventually(t, func() bool { + _, err := host.GetAlarm(t.Context(), ActorType, actor.SingletonActorID, alarmRecurringSync) + return err == nil + }, 10*time.Second, 20*time.Millisecond) + + module.ScheduleSync(t.Context()) + + properties, err := host.GetAlarm(t.Context(), ActorType, actor.SingletonActorID, alarmScheduledSync) + require.NoError(t, err) + assert.WithinDuration(t, time.Now().Add(scheduledSyncDelay), properties.DueTime, time.Second) +} + +// newSyncActorForTest starts a test actor host and allocates the actor without registering it +func newSyncActorForTest(t *testing.T, syncer syncer, scheduleDisabled bool) (*local.Host, *syncActor) { + t.Helper() + + host := testutils.NewActorHostForTest(t, nil) + act, ok := newActor(syncer, scheduleDisabled)(actor.SingletonActorID, host.Service()).(*syncActor) + require.True(t, ok) + + return host, act +} diff --git a/backend/internal/dto/scim_dto.go b/backend/internal/scimsync/dto.go similarity index 73% rename from backend/internal/dto/scim_dto.go rename to backend/internal/scimsync/dto.go index c04c8948..842ba7d5 100644 --- a/backend/internal/dto/scim_dto.go +++ b/backend/internal/scimsync/dto.go @@ -1,18 +1,19 @@ -package dto +package scimsync import ( "time" + "github.com/pocket-id/pocket-id/backend/internal/dto" datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" ) type ScimServiceProviderDTO struct { - ID string `json:"id"` - Endpoint string `json:"endpoint"` - Token string `json:"token"` - LastSyncedAt *datatype.DateTime `json:"lastSyncedAt"` - OidcClient OidcClientMetaDataDto `json:"oidcClient"` - CreatedAt datatype.DateTime `json:"createdAt"` + ID string `json:"id"` + Endpoint string `json:"endpoint"` + Token string `json:"token"` + LastSyncedAt *datatype.DateTime `json:"lastSyncedAt"` + OidcClient dto.OidcClientMetaDataDto `json:"oidcClient"` + CreatedAt datatype.DateTime `json:"createdAt"` } type ScimServiceProviderCreateDTO struct { @@ -58,10 +59,10 @@ type ScimListResponse[T any] struct { } type ScimResourceData struct { - ID string `json:"id,omitempty"` - ExternalID string `json:"externalId,omitempty"` - Schemas []string `json:"schemas"` - Meta ScimResourceMeta `json:"meta,omitempty"` + ID string `json:"id,omitempty"` + ExternalID string `json:"externalId,omitempty"` + Schemas []string `json:"schemas"` + Meta *ScimResourceMeta `json:"meta,omitempty"` } type ScimResourceMeta struct { @@ -85,7 +86,11 @@ func (r ScimResourceData) GetSchemas() []string { } func (r ScimResourceData) GetMeta() ScimResourceMeta { - return r.Meta + if r.Meta == nil { + return ScimResourceMeta{} + } + + return *r.Meta } type ScimResource interface { diff --git a/backend/internal/scimsync/handler.go b/backend/internal/scimsync/handler.go new file mode 100644 index 00000000..9c6d3e5c --- /dev/null +++ b/backend/internal/scimsync/handler.go @@ -0,0 +1,135 @@ +package scimsync + +import ( + "net/http" + + "github.com/gin-gonic/gin" + + "github.com/pocket-id/pocket-id/backend/internal/dto" + "github.com/pocket-id/pocket-id/backend/internal/httpserver" +) + +type handler struct { + service *Service +} + +func newHandler(service *Service) *handler { + return &handler{service: service} +} + +// getServiceProviderByClient godoc +// @Summary Get SCIM service provider +// @Description Get the SCIM service provider configuration for an OIDC client +// @Tags OIDC +// @Produce json +// @Param id path string true "Client ID" +// @Success 200 {object} ScimServiceProviderDTO "SCIM service provider configuration" +// @Failure default {object} dto.ErrorDto "Error" +// @Router /api/oidc/clients/{id}/scim-service-provider [get] +func (h *handler) getServiceProviderByClient(c *gin.Context) error { + provider, err := h.service.GetServiceProviderByClient(c.Request.Context(), c.Param("id")) + if err != nil { + return err + } + + return respondWithServiceProvider(c, http.StatusOK, provider) +} + +// syncServiceProvider godoc +// @Summary Sync SCIM service provider +// @Description Trigger synchronization for a SCIM service provider +// @Tags SCIM +// @Param id path string true "Service Provider ID" +// @Success 200 "OK" +// @Failure default {object} dto.ErrorDto "Error" +// @Router /api/scim/service-provider/{id}/sync [post] +func (h *handler) syncServiceProvider(c *gin.Context) error { + // The sync runs inline rather than through the actor so the response reports whether it succeeded + err := h.service.SyncServiceProvider(c.Request.Context(), c.Param("id")) + if err != nil { + return err + } + + c.Status(http.StatusOK) + return nil +} + +// createServiceProvider godoc +// @Summary Create SCIM service provider +// @Description Create a new SCIM service provider +// @Tags SCIM +// @Accept json +// @Produce json +// @Param serviceProvider body ScimServiceProviderCreateDTO true "SCIM service provider information" +// @Success 201 {object} ScimServiceProviderDTO "Created SCIM service provider" +// @Failure default {object} dto.ErrorDto "Error" +// @Router /api/scim/service-provider [post] +func (h *handler) createServiceProvider(c *gin.Context) error { + var input ScimServiceProviderCreateDTO + err := httpserver.BindJSON(c, &input) + if err != nil { + return err + } + + provider, err := h.service.CreateServiceProvider(c.Request.Context(), &input) + if err != nil { + return err + } + + return respondWithServiceProvider(c, http.StatusCreated, provider) +} + +// updateServiceProvider godoc +// @Summary Update SCIM service provider +// @Description Update an existing SCIM service provider +// @Tags SCIM +// @Accept json +// @Produce json +// @Param id path string true "Service Provider ID" +// @Param serviceProvider body ScimServiceProviderCreateDTO true "SCIM service provider information" +// @Success 200 {object} ScimServiceProviderDTO "Updated SCIM service provider" +// @Failure default {object} dto.ErrorDto "Error" +// @Router /api/scim/service-provider/{id} [put] +func (h *handler) updateServiceProvider(c *gin.Context) error { + var input ScimServiceProviderCreateDTO + err := httpserver.BindJSON(c, &input) + if err != nil { + return err + } + + provider, err := h.service.UpdateServiceProvider(c.Request.Context(), c.Param("id"), &input) + if err != nil { + return err + } + + return respondWithServiceProvider(c, http.StatusOK, provider) +} + +// deleteServiceProvider godoc +// @Summary Delete SCIM service provider +// @Description Delete a SCIM service provider by ID +// @Tags SCIM +// @Param id path string true "Service Provider ID" +// @Success 204 "No Content" +// @Failure default {object} dto.ErrorDto "Error" +// @Router /api/scim/service-provider/{id} [delete] +func (h *handler) deleteServiceProvider(c *gin.Context) error { + err := h.service.DeleteServiceProvider(c.Request.Context(), c.Param("id")) + if err != nil { + return err + } + + c.Status(http.StatusNoContent) + return nil +} + +func respondWithServiceProvider(c *gin.Context, status int, provider ServiceProvider) error { + var output ScimServiceProviderDTO + err := dto.MapStruct(provider, &output) + if err != nil { + return err + } + + c.JSON(status, output) + return nil +} diff --git a/backend/internal/scimsync/model.go b/backend/internal/scimsync/model.go new file mode 100644 index 00000000..d7497609 --- /dev/null +++ b/backend/internal/scimsync/model.go @@ -0,0 +1,21 @@ +package scimsync + +import ( + "github.com/pocket-id/pocket-id/backend/internal/model" + datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" +) + +type ServiceProvider struct { + model.Base + + Endpoint string `sortable:"true"` + Token datatype.EncryptedString + LastSyncedAt *datatype.DateTime `sortable:"true"` + + OidcClientID string + OidcClient model.OidcClient `gorm:"foreignKey:OidcClientID;references:ID;"` +} + +func (ServiceProvider) TableName() string { + return "scim_service_providers" +} diff --git a/backend/internal/scimsync/module.go b/backend/internal/scimsync/module.go new file mode 100644 index 00000000..481aef5b --- /dev/null +++ b/backend/internal/scimsync/module.go @@ -0,0 +1,70 @@ +package scimsync + +import ( + "context" + "fmt" + "log/slog" + "net/http" + + "github.com/gin-gonic/gin" + "github.com/italypaleale/francis/actor" + "github.com/italypaleale/francis/host/local" + "gorm.io/gorm" + + "github.com/pocket-id/pocket-id/backend/internal/httpserver" +) + +type Dependencies struct { + DB *gorm.DB + Actors *local.Host + HTTPClient *http.Client + + // ScheduleDisabled keeps automatic synchronizations from being armed + // It's set in the test environment, where SCIM syncs are driven explicitly by the end-to-end tests + ScheduleDisabled bool +} + +type Module struct { + service *Service + handler *handler + actors *actor.Service + scheduleDisabled bool +} + +func New(deps Dependencies) (*Module, error) { + service := newService(deps.DB, deps.HTTPClient) + + // Register the singleton so recurring and debounced synchronizations run once for the entire cluster + err := deps.Actors.RegisterSingletonActor(ActorType, NewActor(service, deps.ScheduleDisabled)) + if err != nil { + return nil, fmt.Errorf("error registering the %s actor: %w", ActorType, err) + } + + return &Module{ + service: service, + handler: newHandler(service), + actors: deps.Actors.Service(), + scheduleDisabled: deps.ScheduleDisabled, + }, nil +} + +// RegisterRoutes mounts the SCIM service provider endpoints +func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, auth gin.HandlerFunc) { + apiGroup.GET("/oidc/clients/:id/scim-service-provider", auth, httpserver.Handle(m.handler.getServiceProviderByClient)) + apiGroup.POST("/scim/service-provider", auth, httpserver.Handle(m.handler.createServiceProvider)) + apiGroup.POST("/scim/service-provider/:id/sync", auth, httpserver.Handle(m.handler.syncServiceProvider)) + apiGroup.PUT("/scim/service-provider/:id", auth, httpserver.Handle(m.handler.updateServiceProvider)) + apiGroup.DELETE("/scim/service-provider/:id", auth, httpserver.Handle(m.handler.deleteServiceProvider)) +} + +// ScheduleSync schedules a debounced cluster-wide synchronization after SCIM-relevant data changes +func (m *Module) ScheduleSync(ctx context.Context) { + if m.scheduleDisabled { + return + } + + _, err := m.actors.Invoke(ctx, ActorType, actor.SingletonActorID, methodScheduleSync, nil) + if err != nil { + slog.ErrorContext(ctx, "Failed to schedule SCIM sync", slog.Any("error", err)) + } +} diff --git a/backend/internal/service/scim_service.go b/backend/internal/scimsync/service.go similarity index 60% rename from backend/internal/service/scim_service.go rename to backend/internal/scimsync/service.go index 20b04689..89174c35 100644 --- a/backend/internal/service/scim_service.go +++ b/backend/internal/scimsync/service.go @@ -1,8 +1,9 @@ -package service +package scimsync import ( "bytes" "context" + "database/sql" "encoding/json" "errors" "fmt" @@ -14,16 +15,16 @@ import ( "slices" "strconv" "strings" + "sync" "time" - "github.com/go-co-op/gocron/v2" + "gorm.io/gorm" + "github.com/pocket-id/pocket-id/backend/internal/apperror" - "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/model" datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" "github.com/pocket-id/pocket-id/backend/internal/oidc" "github.com/pocket-id/pocket-id/backend/internal/utils" - "gorm.io/gorm" ) const ( @@ -32,7 +33,10 @@ const ( scimContentType = "application/scim+json" ) -const scimErrorBodyLimit = 4096 +const ( + scimErrorBodyLimit = 4 << 10 // 4KB + syncProviderConcurrency = 4 +) type scimSyncAction int @@ -49,113 +53,133 @@ type scimSyncStats struct { Deleted int } -// ScimService handles SCIM provisioning to external service providers. -type ScimService struct { +// Service handles SCIM provisioning to external service providers +type Service struct { db *gorm.DB - scheduler Scheduler httpClient *http.Client } -func NewScimService(db *gorm.DB, scheduler Scheduler, httpClient *http.Client) *ScimService { +func newService(db *gorm.DB, httpClient *http.Client) *Service { if httpClient == nil { - httpClient = &http.Client{Timeout: 20 * time.Second} + httpClient = http.DefaultClient } - return &ScimService{db: db, scheduler: scheduler, httpClient: httpClient} + return &Service{ + db: db, + httpClient: httpClient, + } } -func (s *ScimService) GetServiceProvider( - ctx context.Context, - serviceProviderID string, -) (model.ScimServiceProvider, error) { - var provider model.ScimServiceProvider - err := s.db.WithContext(ctx). +func (s *Service) GetServiceProvider(ctx context.Context, serviceProviderID string) (ServiceProvider, error) { + return getServiceProvider(ctx, s.db, serviceProviderID) +} + +func getServiceProvider(ctx context.Context, db *gorm.DB, serviceProviderID string) (ServiceProvider, error) { + var provider ServiceProvider + err := db.WithContext(ctx). Preload("OidcClient"). Preload("OidcClient.AllowedUserGroups"). First(&provider, "id = ?", serviceProviderID). Error if errors.Is(err, gorm.ErrRecordNotFound) { - return model.ScimServiceProvider{}, apperror.NotFound("SCIM service provider") - } - if err != nil { - return model.ScimServiceProvider{}, err + return ServiceProvider{}, apperror.NotFound("SCIM service provider") + } else if err != nil { + return ServiceProvider{}, err } + return provider, nil } -func (s *ScimService) ListServiceProviders(ctx context.Context) ([]model.ScimServiceProvider, error) { - var providers []model.ScimServiceProvider +func (s *Service) ListServiceProviders(ctx context.Context) ([]ServiceProvider, error) { + var providers []ServiceProvider err := s.db.WithContext(ctx). - Preload("OidcClient"). + Select("id"). Find(&providers). Error if err != nil { return nil, err } + return providers, nil } -func (s *ScimService) CreateServiceProvider( - ctx context.Context, - input *dto.ScimServiceProviderCreateDTO) (model.ScimServiceProvider, error) { - tx := s.db.Begin() - defer func() { - tx.Rollback() - }() - - if err := ensureScimOIDCClientExists(ctx, tx, input.OidcClientID); err != nil { - return model.ScimServiceProvider{}, err - } - - provider := model.ScimServiceProvider{ - Endpoint: input.Endpoint, - Token: datatype.EncryptedString(input.Token), - OidcClientID: input.OidcClientID, - } - - if err := tx.WithContext(ctx).Create(&provider).Error; err != nil { - return model.ScimServiceProvider{}, err - } - if err := tx.Commit().Error; err != nil { - return model.ScimServiceProvider{}, err +func (s *Service) GetServiceProviderByClient(ctx context.Context, clientID string) (ServiceProvider, error) { + var provider ServiceProvider + err := s.db.WithContext(ctx). + First(&provider, "oidc_client_id = ?", clientID). + Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return ServiceProvider{}, apperror.NotFound("SCIM service provider") + } else if err != nil { + return ServiceProvider{}, err } return provider, nil } -func (s *ScimService) UpdateServiceProvider(ctx context.Context, - serviceProviderID string, - input *dto.ScimServiceProviderCreateDTO, -) (model.ScimServiceProvider, error) { +func (s *Service) CreateServiceProvider(ctx context.Context, input *ScimServiceProviderCreateDTO) (ServiceProvider, error) { tx := s.db.Begin() defer func() { tx.Rollback() }() - var provider model.ScimServiceProvider + err := ensureScimOIDCClientExists(ctx, tx, input.OidcClientID) + if err != nil { + return ServiceProvider{}, err + } + + provider := ServiceProvider{ + Endpoint: input.Endpoint, + Token: datatype.EncryptedString(input.Token), + OidcClientID: input.OidcClientID, + } + + err = tx.WithContext(ctx).Create(&provider).Error + if err != nil { + return ServiceProvider{}, fmt.Errorf("error creating service provider: %w", err) + } + + err = tx.Commit().Error + if err != nil { + return ServiceProvider{}, fmt.Errorf("error committing transaction: %w", err) + } + + return provider, nil +} + +func (s *Service) UpdateServiceProvider(ctx context.Context, serviceProviderID string, input *ScimServiceProviderCreateDTO) (ServiceProvider, error) { + tx := s.db.Begin() + defer func() { + tx.Rollback() + }() + + var provider ServiceProvider err := tx.WithContext(ctx). First(&provider, "id = ?", serviceProviderID). Error if errors.Is(err, gorm.ErrRecordNotFound) { - return model.ScimServiceProvider{}, apperror.NotFound("SCIM service provider") - } - if err != nil { - return model.ScimServiceProvider{}, err + return ServiceProvider{}, apperror.NotFound("SCIM service provider") + } else if err != nil { + return ServiceProvider{}, fmt.Errorf("error loading SCIM service provider: %w", err) } - if err := ensureScimOIDCClientExists(ctx, tx, input.OidcClientID); err != nil { - return model.ScimServiceProvider{}, err + err = ensureScimOIDCClientExists(ctx, tx, input.OidcClientID) + if err != nil { + return ServiceProvider{}, err } provider.Endpoint = input.Endpoint provider.Token = datatype.EncryptedString(input.Token) provider.OidcClientID = input.OidcClientID - if err := tx.WithContext(ctx).Save(&provider).Error; err != nil { - return model.ScimServiceProvider{}, err + err = tx.WithContext(ctx).Save(&provider).Error + if err != nil { + return ServiceProvider{}, fmt.Errorf("error saving SCIM service provider: %w", err) } - if err := tx.Commit().Error; err != nil { - return model.ScimServiceProvider{}, err + + err = tx.Commit().Error + if err != nil { + return ServiceProvider{}, fmt.Errorf("error committing transaction: %w", err) } return provider, nil @@ -174,9 +198,10 @@ func ensureScimOIDCClientExists(ctx context.Context, db *gorm.DB, clientID strin return err } -func (s *ScimService) DeleteServiceProvider(ctx context.Context, serviceProviderID string) error { - result := s.db.WithContext(ctx). - Delete(&model.ScimServiceProvider{}, "id = ?", serviceProviderID) +func (s *Service) DeleteServiceProvider(ctx context.Context, serviceProviderID string) error { + result := s.db. + WithContext(ctx). + Delete(&ServiceProvider{}, "id = ?", serviceProviderID) if result.Error != nil { return result.Error } @@ -187,72 +212,36 @@ func (s *ScimService) DeleteServiceProvider(ctx context.Context, serviceProvider return nil } -//nolint:contextcheck -func (s *ScimService) ScheduleSync() { - jobName := "ScheduledScimSync" - start := time.Now().Add(5 * time.Minute) - - _ = s.scheduler.RemoveJob(jobName) - - err := s.scheduler.RegisterJob( - context.Background(), jobName, - gocron.OneTimeJob(gocron.OneTimeJobStartDateTime(start)), s.SyncAll, RegisterJobOpts{}) - - if err != nil { - slog.Error("Failed to schedule SCIM sync", slog.Any("error", err)) - } -} - -func (s *ScimService) SyncAll(ctx context.Context) error { +func (s *Service) SyncAll(ctx context.Context) error { providers, err := s.ListServiceProviders(ctx) if err != nil { return err } - var errs []error - for _, provider := range providers { - if ctx.Err() != nil { - errs = append(errs, ctx.Err()) - break - } - err = s.SyncServiceProvider(ctx, provider.ID) - if err != nil { - errs = append(errs, fmt.Errorf("failed to sync SCIM provider %s: %w", provider.ID, err)) - } - } - return errors.Join(errs...) + return syncServiceProviders(ctx, providers, s.SyncServiceProvider) } -func (s *ScimService) SyncServiceProvider(ctx context.Context, serviceProviderID string) error { +func (s *Service) SyncServiceProvider(ctx context.Context, serviceProviderID string) error { start := time.Now() - provider, err := s.GetServiceProvider(ctx, serviceProviderID) + + // Load one consistent local snapshot and release the transaction before making remote requests + snapshot, err := s.loadSyncSnapshot(ctx, serviceProviderID) if err != nil { return err } + provider := snapshot.provider slog.InfoContext(ctx, "Syncing SCIM service provider", slog.String("provider_id", provider.ID), slog.String("oidc_client_id", provider.OidcClientID), ) - allowedGroupIDs := groupIDs(provider.OidcClient.AllowedUserGroups) - - // Load users and groups that should be synced to the SCIM provider - groups, err := s.groupsForClient(ctx, provider.OidcClient, allowedGroupIDs) - if err != nil { - return err - } - users, err := s.usersForClient(ctx, provider.OidcClient, allowedGroupIDs) - if err != nil { - return err - } - // Load users and groups that already exist in the SCIM provider - userResources, err := listScimResources[dto.ScimUser](s, ctx, provider, "/Users") + userResources, err := listScimResources[ScimUser](s, ctx, provider, "/Users") if err != nil { return err } - groupResources, err := listScimResources[dto.ScimGroup](s, ctx, provider, "/Groups") + groupResources, err := listScimResources[ScimGroup](s, ctx, provider, "/Groups") if err != nil { return err } @@ -260,12 +249,12 @@ func (s *ScimService) SyncServiceProvider(ctx context.Context, serviceProviderID var errs []error // Sync users first, so that groups can reference them - userStats, err := s.syncUsers(ctx, provider, users, &userResources) + userStats, err := s.syncUsers(ctx, provider, snapshot.users, &userResources) if err != nil { errs = append(errs, err) } - groupStats, err := s.syncGroups(ctx, provider, groups, groupResources.Resources, userResources.Resources) + groupStats, err := s.syncGroups(ctx, provider, snapshot.groups, groupResources.Resources, userResources.Resources) if err != nil { errs = append(errs, err) } @@ -287,10 +276,16 @@ func (s *ScimService) SyncServiceProvider(ctx context.Context, serviceProviderID return err } - provider.LastSyncedAt = new(datatype.DateTime(time.Now())) - err = s.db.WithContext(ctx).Save(&provider).Error - if err != nil { - return err + lastSyncedAt := datatype.DateTime(time.Now()) + result := s.db.WithContext(ctx). + Model(&ServiceProvider{}). + Where("id = ?", provider.ID). + Update("last_synced_at", &lastSyncedAt) + if result.Error != nil { + return result.Error + } + if result.RowsAffected == 0 { + return apperror.NotFound("SCIM service provider") } slog.InfoContext(ctx, "SCIM sync completed", @@ -307,12 +302,86 @@ func (s *ScimService) SyncServiceProvider(ctx context.Context, serviceProviderID return nil } -func (s *ScimService) syncUsers( - ctx context.Context, - provider model.ScimServiceProvider, - users []model.User, - resourceList *dto.ScimListResponse[dto.ScimUser], -) (stats scimSyncStats, err error) { +type syncSnapshot struct { + provider ServiceProvider + users []model.User + groups []model.UserGroup +} + +// loadSyncSnapshot reads all local inputs from one point in time without holding the transaction across remote SCIM calls +func (s *Service) loadSyncSnapshot(ctx context.Context, serviceProviderID string) (snapshot syncSnapshot, oErr error) { + oErr = s.db. + WithContext(ctx). + Transaction( + func(tx *gorm.DB) (err error) { + snapshot.provider, err = getServiceProvider(ctx, tx, serviceProviderID) + if err != nil { + return err + } + + allowedGroupIDs := groupIDs(snapshot.provider.OidcClient.AllowedUserGroups) + snapshot.groups, err = groupsForClient(ctx, tx, snapshot.provider.OidcClient, allowedGroupIDs) + if err != nil { + return err + } + + snapshot.users, err = usersForClient(ctx, tx, snapshot.provider.OidcClient, allowedGroupIDs) + return err + }, + syncSnapshotTxOptions(s.db.Name()), + ) + if oErr != nil { + return syncSnapshot{}, oErr + } + + return snapshot, nil +} + +// syncSnapshotTxOptions pins a consistent read snapshot without taking SQLite's configured immediate write lock +func syncSnapshotTxOptions(provider string) *sql.TxOptions { + opts := &sql.TxOptions{ReadOnly: true} + if provider == "postgres" { + opts.Isolation = sql.LevelRepeatableRead + } + + return opts +} + +func syncServiceProviders(ctx context.Context, providers []ServiceProvider, syncProvider func(context.Context, string) error) error { + // Bound concurrency so several independent providers make progress without overwhelming the database or network + semaphore := make(chan struct{}, syncProviderConcurrency) + errs := make([]error, len(providers)) + var waitGroup sync.WaitGroup + + // Start each provider when a slot is available and retain its error in deterministic provider order +providerLoop: + for i, provider := range providers { + select { + case semaphore <- struct{}{}: + case <-ctx.Done(): + errs[i] = ctx.Err() + break providerLoop + } + + waitGroup.Go(func() { + defer func() { + <-semaphore + }() + + err := syncProvider(ctx, provider.ID) + if err != nil { + errs[i] = fmt.Errorf("failed to sync SCIM provider %s: %w", provider.ID, err) + } + }) + } + + // Wait for every started provider so one failure never prevents the remaining providers from synchronizing + waitGroup.Wait() + + return errors.Join(errs...) +} + +func (s *Service) syncUsers(ctx context.Context, provider ServiceProvider, users []model.User, resourceList *ScimListResponse[ScimUser]) (stats scimSyncStats, err error) { var errs []error // Update or create users @@ -340,7 +409,7 @@ func (s *ScimService) syncUsers( } } - // Delete users that are present in SCIM provider but not locally. + // Delete users that are present in SCIM provider but not locally userSet := make(map[string]struct{}) for _, u := range users { userSet[u.ID] = struct{}{} @@ -359,13 +428,7 @@ func (s *ScimService) syncUsers( return stats, errors.Join(errs...) } -func (s *ScimService) syncGroups( - ctx context.Context, - provider model.ScimServiceProvider, - groups []model.UserGroup, - remoteGroups []dto.ScimGroup, - userResources []dto.ScimUser, -) (stats scimSyncStats, err error) { +func (s *Service) syncGroups(ctx context.Context, provider ServiceProvider, groups []model.UserGroup, remoteGroups []ScimGroup, userResources []ScimUser) (stats scimSyncStats, err error) { var errs []error // Update or create groups @@ -410,23 +473,19 @@ func (s *ScimService) syncGroups( return stats, errors.Join(errs...) } -func (s *ScimService) syncUser(ctx context.Context, - provider model.ScimServiceProvider, - user model.User, - userResource *dto.ScimUser, -) (scimSyncAction, *dto.ScimUser, error) { +func (s *Service) syncUser(ctx context.Context, provider ServiceProvider, user model.User, userResource *ScimUser) (scimSyncAction, *ScimUser, error) { // If user is not allowed for the client, delete it from SCIM provider if userResource != nil && !oidc.IsUserGroupAllowedToAuthorize(user, provider.OidcClient) { return scimActionDeleted, nil, s.deleteScimResource(ctx, provider, fmt.Sprintf("/Users/%s", url.PathEscape(userResource.ID))) } - payload := dto.ScimUser{ - ScimResourceData: dto.ScimResourceData{ + payload := ScimUser{ + ScimResourceData: ScimResourceData{ Schemas: []string{scimUserSchema}, ExternalID: user.ID, }, UserName: user.Username, - Name: &dto.ScimName{ + Name: &ScimName{ GivenName: user.FirstName, FamilyName: user.LastName, }, @@ -435,7 +494,7 @@ func (s *ScimService) syncUser(ctx context.Context, } if user.Email != nil { - payload.Emails = []dto.ScimEmail{{ + payload.Emails = []ScimEmail{{ Value: *user.Email, Primary: true, }} @@ -463,20 +522,18 @@ func (s *ScimService) syncUser(ctx context.Context, return scimActionCreated, userResource, nil } -func (s *ScimService) syncGroup( - ctx context.Context, - provider model.ScimServiceProvider, - group model.UserGroup, - groupResource *dto.ScimGroup, - userResources []dto.ScimUser, -) (scimSyncAction, error) { +func (s *Service) syncGroup(ctx context.Context, provider ServiceProvider, group model.UserGroup, groupResource *ScimGroup, userResources []ScimUser) (scimSyncAction, error) { // If group is not allowed for the client, delete it from SCIM provider if groupResource != nil && !groupAllowedForClient(group.ID, provider.OidcClient) { - return scimActionDeleted, s.deleteScimResource(ctx, provider, fmt.Sprintf("/Groups/%s", url.PathEscape(groupResource.GetID()))) + err := s.deleteScimResource(ctx, provider, fmt.Sprintf("/Groups/%s", url.PathEscape(groupResource.GetID()))) + if err != nil { + return scimActionNone, err + } + return scimActionDeleted, nil } // Prepare group members - members := make([]dto.ScimGroupMember, len(group.Users)) + members := make([]ScimGroupMember, len(group.Users)) for i, user := range group.Users { userResource := getResourceByExternalID(user.ID, userResources) if userResource == nil { @@ -484,13 +541,13 @@ func (s *ScimService) syncGroup( return scimActionNone, fmt.Errorf("cannot sync group %s: user %s is not provisioned in SCIM provider", group.ID, user.ID) } - members[i] = dto.ScimGroupMember{ + members[i] = ScimGroupMember{ Value: userResource.GetID(), } } - groupPayload := dto.ScimGroup{ - ScimResourceData: dto.ScimResourceData{ + groupPayload := ScimGroup{ + ScimResourceData: ScimResourceData{ Schemas: []string{scimGroupSchema}, ExternalID: group.ID, }, @@ -542,14 +599,10 @@ func groupIDs(groups []model.UserGroup) []string { return ids } -func (s *ScimService) groupsForClient( - ctx context.Context, - client model.OidcClient, - allowedGroupIDs []string, -) ([]model.UserGroup, error) { +func groupsForClient(ctx context.Context, db *gorm.DB, client model.OidcClient, allowedGroupIDs []string) ([]model.UserGroup, error) { var groups []model.UserGroup - query := s.db.WithContext(ctx).Preload("Users").Model(&model.UserGroup{}) + query := db.WithContext(ctx).Preload("Users").Model(&model.UserGroup{}) if client.IsGroupRestricted { if len(allowedGroupIDs) == 0 { return groups, nil @@ -557,20 +610,17 @@ func (s *ScimService) groupsForClient( query = query.Where("id IN ?", allowedGroupIDs) } - if err := query.Find(&groups).Error; err != nil { + err := query.Find(&groups).Error + if err != nil { return nil, err } return groups, nil } -func (s *ScimService) usersForClient( - ctx context.Context, - client model.OidcClient, - allowedGroupIDs []string, -) ([]model.User, error) { +func usersForClient(ctx context.Context, db *gorm.DB, client model.OidcClient, allowedGroupIDs []string) ([]model.User, error) { var users []model.User - query := s.db.WithContext(ctx).Model(&model.User{}) + query := db.WithContext(ctx).Model(&model.User{}) if client.IsGroupRestricted { if len(allowedGroupIDs) == 0 { return users, nil @@ -584,13 +634,14 @@ func (s *ScimService) usersForClient( query = query.Preload("UserGroups") - if err := query.Find(&users).Error; err != nil { + err := query.Find(&users).Error + if err != nil { return nil, err } return users, nil } -func getResourceByExternalID[T dto.ScimResource](externalID string, resource []T) *T { +func getResourceByExternalID[T ScimResource](externalID string, resource []T) *T { for i := range resource { if resource[i].GetExternalID() == externalID { return &resource[i] @@ -599,12 +650,7 @@ func getResourceByExternalID[T dto.ScimResource](externalID string, resource []T return nil } -func listScimResources[T any]( - s *ScimService, - ctx context.Context, - provider model.ScimServiceProvider, - path string, -) (result dto.ScimListResponse[T], err error) { +func listScimResources[T any](s *Service, ctx context.Context, provider ServiceProvider, path string) (result ScimListResponse[T], err error) { startIndex := 1 count := 1000 @@ -617,16 +663,18 @@ func listScimResources[T any]( resp, err := s.scimRequest(ctx, provider, http.MethodGet, path, nil, queryParams) if err != nil { - return dto.ScimListResponse[T]{}, err + return ScimListResponse[T]{}, err } - if err := ensureScimStatus(ctx, resp, provider, http.StatusOK); err != nil { - return dto.ScimListResponse[T]{}, err + err = ensureScimStatus(ctx, resp, provider, http.StatusOK) + if err != nil { + return ScimListResponse[T]{}, err } - var page dto.ScimListResponse[T] - if err := json.NewDecoder(resp.Body).Decode(&page); err != nil { - return dto.ScimListResponse[T]{}, fmt.Errorf("failed to decode SCIM list response: %w", err) + var page ScimListResponse[T] + err = json.NewDecoder(resp.Body).Decode(&page) + if err != nil { + return ScimListResponse[T]{}, fmt.Errorf("failed to decode SCIM list response: %w", err) } resp.Body.Close() @@ -650,55 +698,49 @@ func listScimResources[T any]( return result, nil } -func createScimResource[T dto.ScimResource]( - s *ScimService, - ctx context.Context, - provider model.ScimServiceProvider, - path string, payload T) (*T, error) { +func createScimResource[T ScimResource](s *Service, ctx context.Context, provider ServiceProvider, path string, payload T) (*T, error) { resp, err := s.scimRequest(ctx, provider, http.MethodPost, path, payload, nil) if err != nil { return nil, err } defer resp.Body.Close() - if err := ensureScimStatus(ctx, resp, provider, http.StatusOK, http.StatusCreated); err != nil { + err = ensureScimStatus(ctx, resp, provider, http.StatusOK, http.StatusCreated) + if err != nil { return nil, err } var resource T - if err := json.NewDecoder(resp.Body).Decode(&resource); err != nil { + err = json.NewDecoder(resp.Body).Decode(&resource) + if err != nil { return nil, fmt.Errorf("failed to decode SCIM create response: %w", err) } return &resource, nil } -func updateScimResource[T dto.ScimResource]( - s *ScimService, - ctx context.Context, - provider model.ScimServiceProvider, - path string, - payload T, -) (*T, error) { +func updateScimResource[T ScimResource](s *Service, ctx context.Context, provider ServiceProvider, path string, payload T) (*T, error) { resp, err := s.scimRequest(ctx, provider, http.MethodPut, path, payload, nil) if err != nil { return nil, err } defer resp.Body.Close() - if err := ensureScimStatus(ctx, resp, provider, http.StatusOK, http.StatusCreated); err != nil { + err = ensureScimStatus(ctx, resp, provider, http.StatusOK, http.StatusCreated) + if err != nil { return nil, err } var resource T - if err := json.NewDecoder(resp.Body).Decode(&resource); err != nil { + err = json.NewDecoder(resp.Body).Decode(&resource) + if err != nil { return nil, fmt.Errorf("failed to decode SCIM update response: %w", err) } return &resource, nil } -func (s *ScimService) deleteScimResource(ctx context.Context, provider model.ScimServiceProvider, path string) error { +func (s *Service) deleteScimResource(ctx context.Context, provider ServiceProvider, path string) error { resp, err := s.scimRequest(ctx, provider, http.MethodDelete, path, nil, nil) if err != nil { return err @@ -709,17 +751,14 @@ func (s *ScimService) deleteScimResource(ctx context.Context, provider model.Sci return nil } - return ensureScimStatus(ctx, resp, provider, http.StatusOK, http.StatusNoContent) + err = ensureScimStatus(ctx, resp, provider, http.StatusOK, http.StatusNoContent) + if err != nil { + return err + } + return nil } -func (s *ScimService) scimRequest( - ctx context.Context, - provider model.ScimServiceProvider, - method, - path string, - payload any, - queryParams map[string]string, -) (*http.Response, error) { +func (s *Service) scimRequest(ctx context.Context, provider ServiceProvider, method, path string, payload any, queryParams map[string]string) (*http.Response, error) { urlString, err := scimURL(provider.Endpoint, path, queryParams) if err != nil { return nil, err @@ -781,7 +820,8 @@ func (s *ScimService) scimRequest( ) resp.Body.Close() - if err := utils.SleepWithContext(ctx, retryDelay); err != nil { + err = utils.SleepWithContext(ctx, retryDelay) + if err != nil { return nil, err } } @@ -792,11 +832,14 @@ func (s *ScimService) scimRequest( func scimRetryDelay(retryAfter string, attempt int) time.Duration { // Respect Retry-After when provided if retryAfter != "" { - if seconds, err := strconv.Atoi(retryAfter); err == nil { + seconds, err := strconv.Atoi(retryAfter) + if err == nil { return time.Duration(seconds) * time.Second } - if t, err := http.ParseTime(retryAfter); err == nil { - if delay := time.Until(t); delay > 0 { + t, err := http.ParseTime(retryAfter) + if err == nil { + delay := time.Until(t) + if delay > 0 { return delay } } @@ -828,11 +871,7 @@ func scimURL(endpoint, p string, queryParams map[string]string) (string, error) return u.String(), nil } -func ensureScimStatus( - ctx context.Context, - resp *http.Response, - provider model.ScimServiceProvider, - allowedStatuses ...int) error { +func ensureScimStatus(ctx context.Context, resp *http.Response, provider ServiceProvider, allowedStatuses ...int) error { if slices.Contains(allowedStatuses, resp.StatusCode) { return nil } diff --git a/backend/internal/scimsync/service_test.go b/backend/internal/scimsync/service_test.go new file mode 100644 index 00000000..e112b806 --- /dev/null +++ b/backend/internal/scimsync/service_test.go @@ -0,0 +1,136 @@ +package scimsync + +import ( + "context" + "database/sql" + "errors" + "fmt" + "sync/atomic" + "testing" + "time" + + "github.com/pocket-id/pocket-id/backend/internal/apperror" + "github.com/pocket-id/pocket-id/backend/internal/model" + testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing" + "github.com/stretchr/testify/require" +) + +func TestServiceProviderOperationsReturnSpecificNotFoundErrors(t *testing.T) { + service := newService(testutils.NewDatabaseForTest(t), nil) + + _, err := service.CreateServiceProvider(t.Context(), &ScimServiceProviderCreateDTO{ + Endpoint: "https://scim.example.com", + OidcClientID: "missing-client", + }) + require.True(t, apperror.IsCode(err, apperror.CodeNotFound)) + + _, err = service.GetServiceProvider(t.Context(), "missing-provider") + require.True(t, apperror.IsCode(err, apperror.CodeNotFound)) + + err = service.DeleteServiceProvider(t.Context(), "missing-provider") + require.True(t, apperror.IsCode(err, apperror.CodeNotFound)) +} + +func TestServiceProviderCreateAndUpdate(t *testing.T) { + db := testutils.NewDatabaseForTest(t) + service := newService(db, nil) + + // Create two clients so provider creation and reassignment both satisfy the foreign key + require.NoError(t, db.Create(&[]model.OidcClient{ + {Base: model.Base{ID: "client-1"}, Name: "Client 1"}, + {Base: model.Base{ID: "client-2"}, Name: "Client 2"}, + }).Error) + + // Create the provider with its initial client in one transaction + provider, err := service.CreateServiceProvider(t.Context(), &ScimServiceProviderCreateDTO{ + Endpoint: "https://scim.example.com/v1", + OidcClientID: "client-1", + }) + require.NoError(t, err) + require.NotEmpty(t, provider.ID) + + // Move the provider to the second client in one transaction + provider, err = service.UpdateServiceProvider(t.Context(), provider.ID, &ScimServiceProviderCreateDTO{ + Endpoint: "https://scim.example.com/v2", + OidcClientID: "client-2", + }) + require.NoError(t, err) + require.Equal(t, "https://scim.example.com/v2", provider.Endpoint) + require.Equal(t, "client-2", provider.OidcClientID) + + // Verify the committed provider retains both updated values + persisted, err := service.GetServiceProvider(t.Context(), provider.ID) + require.NoError(t, err) + require.Equal(t, provider.Endpoint, persisted.Endpoint) + require.Equal(t, provider.OidcClientID, persisted.OidcClientID) + + // Verify SQLite accepts the read-only snapshot used by synchronization + snapshot, err := service.loadSyncSnapshot(t.Context(), provider.ID) + require.NoError(t, err) + require.Equal(t, provider.ID, snapshot.provider.ID) +} + +func TestSyncSnapshotTxOptions(t *testing.T) { + require.Equal(t, &sql.TxOptions{Isolation: sql.LevelRepeatableRead, ReadOnly: true}, syncSnapshotTxOptions("postgres")) + require.Equal(t, &sql.TxOptions{ReadOnly: true}, syncSnapshotTxOptions("sqlite")) +} + +func TestSyncServiceProvidersLimitsConcurrencyAndJoinsErrors(t *testing.T) { + providers := make([]ServiceProvider, 8) + for i := range providers { + providers[i].ID = fmt.Sprintf("provider-%d", i) + } + + started := make(chan string, len(providers)) + release := make(chan struct{}) + done := make(chan error, 1) + var active atomic.Int32 + var maximum atomic.Int32 + + go func() { + done <- syncServiceProviders(t.Context(), providers, func(_ context.Context, providerID string) error { + current := active.Add(1) + for { + previous := maximum.Load() + if current <= previous || maximum.CompareAndSwap(previous, current) { + break + } + } + started <- providerID + <-release + active.Add(-1) + + if providerID == "provider-0" || providerID == "provider-7" { + return errors.New(providerID + " failed") + } + return nil + }) + }() + + // Keep the first batch blocked so a fifth provider would expose a broken concurrency limit + firstBatch := make([]string, 0, syncProviderConcurrency) + +firstBatchLoop: + for range syncProviderConcurrency { + select { + case providerID := <-started: + firstBatch = append(firstBatch, providerID) + case <-time.After(2 * time.Second): + break firstBatchLoop + } + } + + var unexpectedProvider string + select { + case unexpectedProvider = <-started: + case <-time.After(100 * time.Millisecond): + } + close(release) + + err := <-done + require.Len(t, firstBatch, syncProviderConcurrency) + require.Empty(t, unexpectedProvider) + require.EqualValues(t, syncProviderConcurrency, maximum.Load()) + require.ErrorContains(t, err, "provider-0 failed") + require.ErrorContains(t, err, "provider-7 failed") +} diff --git a/backend/internal/scimsync/sync_test.go b/backend/internal/scimsync/sync_test.go new file mode 100644 index 00000000..44870be6 --- /dev/null +++ b/backend/internal/scimsync/sync_test.go @@ -0,0 +1,837 @@ +package scimsync + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "maps" + "net/http" + "net/url" + "slices" + "sort" + "strconv" + "strings" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" + + "github.com/pocket-id/pocket-id/backend/internal/model" + datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" + testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing" +) + +const ( + mockSCIMEndpoint = "https://scim.example.test" + scimListResponseSchema = "urn:ietf:params:scim:api:messages:2.0:ListResponse" + scimErrorResponseSchema = "urn:ietf:params:scim:api:messages:2.0:Error" + mockSCIMRequestContentType = "application/scim+json" +) + +type scimSyncFixture struct { + db *gorm.DB + service *Service + transport *mockSCIMTransport + client model.OidcClient + provider ServiceProvider +} + +func newSCIMSyncFixture(t *testing.T, restricted bool) *scimSyncFixture { + t.Helper() + + db := testutils.NewDatabaseForTest(t) + providerToken := t.Name() + client := model.OidcClient{ + Base: model.Base{ID: "oidc-client"}, + Name: "SCIM client", + IsGroupRestricted: restricted, + } + err := db.Create(&client).Error + require.NoError(t, err) + + provider := ServiceProvider{ + Base: model.Base{ID: "scim-provider"}, + Endpoint: mockSCIMEndpoint, + Token: datatype.EncryptedString(providerToken), + OidcClientID: client.ID, + } + err = db.Create(&provider).Error + require.NoError(t, err) + + transport := newMockSCIMTransport(providerToken) + service := newService(db, &http.Client{Transport: transport}) + + return &scimSyncFixture{ + db: db, + service: service, + transport: transport, + client: client, + provider: provider, + } +} + +func (f *scimSyncFixture) createUser(t *testing.T, id, username string, email *string, disabled bool) model.User { + t.Helper() + + user := model.User{ + Base: model.Base{ID: id}, + Username: username, + Email: email, + FirstName: strings.ToUpper(username[:1]) + username[1:], + LastName: "Example", + DisplayName: username + " display", + Disabled: disabled, + } + err := f.db.Create(&user).Error + require.NoError(t, err) + return user +} + +func (f *scimSyncFixture) createGroup(t *testing.T, id, name string, users ...model.User) model.UserGroup { + t.Helper() + + group := model.UserGroup{ + Base: model.Base{ID: id}, + Name: name, + FriendlyName: name + " friendly", + } + err := f.db.Create(&group).Error + require.NoError(t, err) + if len(users) > 0 { + err = f.db.Model(&group).Association("Users").Replace(users) + require.NoError(t, err) + } + return group +} + +func (f *scimSyncFixture) allowGroups(t *testing.T, groups ...model.UserGroup) { + t.Helper() + err := f.db.Model(&f.client).Association("AllowedUserGroups").Replace(groups) + require.NoError(t, err) +} + +func (f *scimSyncFixture) requireLastSynced(t *testing.T, expected bool) { + t.Helper() + + var provider ServiceProvider + err := f.db.First(&provider, "id = ?", f.provider.ID).Error + require.NoError(t, err) + if expected { + require.NotNil(t, provider.LastSyncedAt) + } else { + require.Nil(t, provider.LastSyncedAt) + } +} + +func TestSyncCreatesCompliantUsersBeforeGroups(t *testing.T) { + fixture := newSCIMSyncFixture(t, false) + aliceEmail := "alice@example.com" + alice := fixture.createUser(t, "user-alice", "alice", &aliceEmail, false) + bob := fixture.createUser(t, "user-bob", "bob", nil, true) + group := fixture.createGroup(t, "group-engineering", "engineering", alice, bob) + + require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID)) + + users := fixture.transport.usersSnapshot() + require.Len(t, users, 2) + remoteAlice := resourceByExternalID(alice.ID, users) + require.NotNil(t, remoteAlice) + assert.Equal(t, "alice", remoteAlice.UserName) + assert.Equal(t, "Alice", remoteAlice.Name.GivenName) + assert.Equal(t, "Example", remoteAlice.Name.FamilyName) + assert.Equal(t, alice.DisplayName, remoteAlice.Display) + assert.True(t, remoteAlice.Active) + require.Equal(t, []ScimEmail{{Value: aliceEmail, Primary: true}}, remoteAlice.Emails) + + remoteBob := resourceByExternalID(bob.ID, users) + require.NotNil(t, remoteBob) + assert.False(t, remoteBob.Active) + assert.Empty(t, remoteBob.Emails) + + groups := fixture.transport.groupsSnapshot() + require.Len(t, groups, 1) + remoteGroup := resourceByExternalID(group.ID, groups) + require.NotNil(t, remoteGroup) + assert.Equal(t, group.FriendlyName, remoteGroup.Display) + assert.ElementsMatch(t, []ScimGroupMember{{Value: remoteAlice.ID}, {Value: remoteBob.ID}}, remoteGroup.Members) + + requests := fixture.transport.requestsSnapshot() + lastUserCreate := lastRequestIndex(requests, http.MethodPost, "/Users") + firstGroupCreate := firstRequestIndex(requests, http.MethodPost, "/Groups") + require.NotEqual(t, -1, lastUserCreate) + require.Greater(t, firstGroupCreate, lastUserCreate) + fixture.requireLastSynced(t, true) + fixture.transport.requireCompliant(t) +} + +func TestSyncUpdatesExistingUsersAndGroupsWithPUT(t *testing.T) { + fixture := newSCIMSyncFixture(t, false) + email := "updated@example.com" + user := fixture.createUser(t, "user-updated", "updated", &email, true) + group := fixture.createGroup(t, "group-updated", "updated", user) + remoteModified := time.Now().Add(-time.Hour) + + fixture.transport.seedUser(ScimUser{ + ScimResourceData: remoteResourceData("remote-user", user.ID, scimUserSchema, "User", remoteModified), + UserName: "stale-name", + Active: true, + }) + fixture.transport.seedGroup(ScimGroup{ + ScimResourceData: remoteResourceData("remote-group", group.ID, scimGroupSchema, "Group", remoteModified), + Display: "stale-group", + }) + + require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID)) + + remoteUser := fixture.transport.user("remote-user") + require.NotNil(t, remoteUser) + assert.Equal(t, user.Username, remoteUser.UserName) + assert.Equal(t, user.DisplayName, remoteUser.Display) + assert.False(t, remoteUser.Active) + require.Equal(t, []ScimEmail{{Value: email, Primary: true}}, remoteUser.Emails) + + remoteGroup := fixture.transport.group("remote-group") + require.NotNil(t, remoteGroup) + assert.Equal(t, group.FriendlyName, remoteGroup.Display) + require.Equal(t, []ScimGroupMember{{Value: remoteUser.ID}}, remoteGroup.Members) + + requests := fixture.transport.requestsSnapshot() + assert.Equal(t, 1, countRequests(requests, http.MethodPut, "/Users/remote-user")) + assert.Equal(t, 1, countRequests(requests, http.MethodPut, "/Groups/remote-group")) + assert.Zero(t, countRequestsWithPrefix(requests, http.MethodPost, "/")) + fixture.requireLastSynced(t, true) + fixture.transport.requireCompliant(t) +} + +func TestSyncRestrictedClientDeletesDisallowedResources(t *testing.T) { + fixture := newSCIMSyncFixture(t, true) + allowedUser := fixture.createUser(t, "user-allowed", "allowed", nil, false) + deniedUser := fixture.createUser(t, "user-denied", "denied", nil, false) + allowedGroup := fixture.createGroup(t, "group-allowed", "allowed", allowedUser) + deniedGroup := fixture.createGroup(t, "group-denied", "denied", deniedUser) + fixture.allowGroups(t, allowedGroup) + remoteModified := time.Now().Add(time.Hour) + + fixture.transport.seedUser(ScimUser{ + ScimResourceData: remoteResourceData("remote-allowed-user", allowedUser.ID, scimUserSchema, "User", remoteModified), + UserName: allowedUser.Username, + Active: true, + }) + fixture.transport.seedUser(ScimUser{ + ScimResourceData: remoteResourceData("remote-denied-user", deniedUser.ID, scimUserSchema, "User", remoteModified), + UserName: deniedUser.Username, + Active: true, + }) + fixture.transport.seedGroup(ScimGroup{ + ScimResourceData: remoteResourceData("remote-allowed-group", allowedGroup.ID, scimGroupSchema, "Group", remoteModified), + Display: allowedGroup.FriendlyName, + Members: []ScimGroupMember{{Value: "remote-allowed-user"}}, + }) + fixture.transport.seedGroup(ScimGroup{ + ScimResourceData: remoteResourceData("remote-denied-group", deniedGroup.ID, scimGroupSchema, "Group", remoteModified), + Display: deniedGroup.FriendlyName, + Members: []ScimGroupMember{{Value: "remote-denied-user"}}, + }) + + require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID)) + + assert.NotNil(t, fixture.transport.user("remote-allowed-user")) + assert.Nil(t, fixture.transport.user("remote-denied-user")) + assert.NotNil(t, fixture.transport.group("remote-allowed-group")) + assert.Nil(t, fixture.transport.group("remote-denied-group")) + requests := fixture.transport.requestsSnapshot() + assert.Equal(t, 1, countRequests(requests, http.MethodDelete, "/Users/remote-denied-user")) + assert.Equal(t, 1, countRequests(requests, http.MethodDelete, "/Groups/remote-denied-group")) + fixture.requireLastSynced(t, true) + fixture.transport.requireCompliant(t) +} + +func TestSyncSkipsResourcesNewerThanTheLocalSnapshotAndPaginates(t *testing.T) { + fixture := newSCIMSyncFixture(t, false) + alice := fixture.createUser(t, "user-alice", "alice", nil, false) + bob := fixture.createUser(t, "user-bob", "bob", nil, false) + remoteModified := time.Now().Add(time.Hour) + fixture.transport.pageSize = 1 + + fixture.transport.seedUser(ScimUser{ + ScimResourceData: remoteResourceData("remote-alice", alice.ID, scimUserSchema, "User", remoteModified), + UserName: alice.Username, + Active: true, + }) + fixture.transport.seedUser(ScimUser{ + ScimResourceData: remoteResourceData("remote-bob", bob.ID, scimUserSchema, "User", remoteModified), + UserName: bob.Username, + Active: true, + }) + + require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID)) + + requests := fixture.transport.requestsSnapshot() + assert.Equal(t, []string{"1", "2"}, queryValues(requests, http.MethodGet, "/Users", "startIndex")) + assert.Equal(t, []string{"1000", "1000"}, queryValues(requests, http.MethodGet, "/Users", "count")) + assert.Zero(t, countMutationRequests(requests)) + fixture.requireLastSynced(t, true) + fixture.transport.requireCompliant(t) +} + +func TestSyncContinuesAfterResourceFailureAndDoesNotMarkCompletion(t *testing.T) { + fixture := newSCIMSyncFixture(t, false) + fixture.createUser(t, "user-success", "success", nil, false) + fixture.createUser(t, "user-failure", "failure", nil, false) + fixture.transport.failCreates["user-failure"] = http.StatusInternalServerError + + err := fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID) + require.Error(t, err) + require.ErrorContains(t, err, "status 500") + + users := fixture.transport.usersSnapshot() + assert.NotNil(t, resourceByExternalID("user-success", users)) + assert.Nil(t, resourceByExternalID("user-failure", users)) + fixture.requireLastSynced(t, false) + fixture.transport.requireCompliant(t) +} + +func TestSyncRetriesRateLimitedSCIMRequests(t *testing.T) { + fixture := newSCIMSyncFixture(t, false) + fixture.transport.rateLimits[http.MethodGet+" /Users"] = 2 + + require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID)) + + requests := fixture.transport.requestsSnapshot() + assert.Equal(t, 3, countRequests(requests, http.MethodGet, "/Users")) + fixture.requireLastSynced(t, true) + fixture.transport.requireCompliant(t) +} + +func remoteResourceData(id, externalID, schema, resourceType string, modified time.Time) ScimResourceData { + return ScimResourceData{ + ID: id, + ExternalID: externalID, + Schemas: []string{schema}, + Meta: &ScimResourceMeta{ + Location: mockSCIMEndpoint + "/" + resourceType + "s/" + id, + ResourceType: resourceType, + Created: modified.Add(-time.Hour), + LastModified: modified, + Version: `W/"seed"`, + }, + } +} + +func resourceByExternalID[T ScimResource](externalID string, resources map[string]T) *T { + for _, resource := range resources { + if resource.GetExternalID() == externalID { + result := resource + return &result + } + } + return nil +} + +type mockSCIMRequest struct { + method string + path string + query url.Values + body []byte +} + +// mockSCIMTransport is a *http.Transport that mocks HTTP server responses to test for compliance with SCIM specs +type mockSCIMTransport struct { + mu sync.Mutex + + expectedToken string + pageSize int + nextID int + users map[string]ScimUser + groups map[string]ScimGroup + requests []mockSCIMRequest + violations []string + rateLimits map[string]int + failCreates map[string]int +} + +func newMockSCIMTransport(expectedToken string) *mockSCIMTransport { + return &mockSCIMTransport{ + expectedToken: expectedToken, + users: map[string]ScimUser{}, + groups: map[string]ScimGroup{}, + rateLimits: map[string]int{}, + failCreates: map[string]int{}, + } +} + +func (m *mockSCIMTransport) RoundTrip(req *http.Request) (*http.Response, error) { + var body []byte + if req.Body != nil { + var err error + body, err = io.ReadAll(req.Body) + if err != nil { + return nil, err + } + } + + m.mu.Lock() + defer m.mu.Unlock() + + m.requests = append(m.requests, mockSCIMRequest{ + method: req.Method, + path: req.URL.Path, + query: req.URL.Query(), + body: slices.Clone(body), + }) + + if violation := m.validateRequest(req, body); violation != "" { + m.violations = append(m.violations, violation) + return mockSCIMErrorResponse(req, http.StatusBadRequest, violation), nil + } + + requestKey := req.Method + " " + req.URL.Path + if m.rateLimits[requestKey] > 0 { + m.rateLimits[requestKey]-- + response := mockSCIMErrorResponse(req, http.StatusTooManyRequests, "rate limited") + response.Header.Set("Retry-After", "0") + return response, nil + } + + segments := strings.Split(strings.Trim(req.URL.Path, "/"), "/") + if len(segments) == 0 || segments[0] == "" { + return mockSCIMErrorResponse(req, http.StatusNotFound, "resource path is empty"), nil + } + + switch segments[0] { + case "Users": + return m.handleUsers(req, segments, body), nil + case "Groups": + return m.handleGroups(req, segments, body), nil + default: + return mockSCIMErrorResponse(req, http.StatusNotFound, "resource type is unknown"), nil + } +} + +func (m *mockSCIMTransport) validateRequest(req *http.Request, body []byte) string { + if req.URL.Scheme != "https" || req.URL.Host != "scim.example.test" { + return fmt.Sprintf("request used unexpected SCIM endpoint %s", req.URL.String()) + } + if req.Header.Get("Accept") != mockSCIMRequestContentType { + return "request did not accept application/scim+json" + } + if req.Header.Get("Authorization") != "Bearer "+m.expectedToken { + return "request did not use the configured bearer token" + } + if len(body) > 0 && req.Header.Get("Content-Type") != mockSCIMRequestContentType { + return "request body did not use application/scim+json" + } + + return "" +} + +func (m *mockSCIMTransport) handleUsers(req *http.Request, segments []string, body []byte) *http.Response { + switch req.Method { + case http.MethodGet: + if len(segments) != 1 { + return mockSCIMErrorResponse(req, http.StatusMethodNotAllowed, "individual user reads are unsupported") + } + resources := make([]ScimUser, 0, len(m.users)) + for _, user := range m.users { + resources = append(resources, user) + } + sort.Slice(resources, func(i, j int) bool { return resources[i].ID < resources[j].ID }) + return mockSCIMListResponse(req, resources, m.pageSize) + + case http.MethodPost: + if len(segments) != 1 { + return mockSCIMErrorResponse(req, http.StatusNotFound, "user collection path is invalid") + } + var user ScimUser + err := json.Unmarshal(body, &user) + if err != nil { + return mockSCIMErrorResponse(req, http.StatusBadRequest, "user payload is invalid JSON") + } + violation := validateUserPayload(body, user) + if violation != "" { + m.violations = append(m.violations, violation) + return mockSCIMErrorResponse(req, http.StatusBadRequest, violation) + } + status := m.failCreates[user.ExternalID] + if status != 0 { + return mockSCIMErrorResponse(req, status, "injected user creation failure") + } + m.nextID++ + user.ID = fmt.Sprintf("remote-user-%d", m.nextID) + user.Meta = newMockMeta("User", "/Users/"+user.ID, m.nextID) + m.users[user.ID] = user + return mockSCIMJSONResponse(req, http.StatusCreated, user) + + case http.MethodPut: + if len(segments) != 2 { + return mockSCIMErrorResponse(req, http.StatusNotFound, "user resource path is invalid") + } + _, ok := m.users[segments[1]] + if !ok { + return mockSCIMErrorResponse(req, http.StatusNotFound, "user does not exist") + } + var user ScimUser + err := json.Unmarshal(body, &user) + if err != nil { + return mockSCIMErrorResponse(req, http.StatusBadRequest, "user payload is invalid JSON") + } + violation := validateUserPayload(body, user) + if violation != "" { + m.violations = append(m.violations, violation) + return mockSCIMErrorResponse(req, http.StatusBadRequest, violation) + } + m.nextID++ + user.ID = segments[1] + user.Meta = newMockMeta("User", "/Users/"+user.ID, m.nextID) + m.users[user.ID] = user + return mockSCIMJSONResponse(req, http.StatusOK, user) + + case http.MethodDelete: + if len(segments) != 2 { + return mockSCIMErrorResponse(req, http.StatusNotFound, "user resource path is invalid") + } + _, ok := m.users[segments[1]] + if !ok { + return mockSCIMErrorResponse(req, http.StatusNotFound, "user does not exist") + } + delete(m.users, segments[1]) + return mockSCIMNoContentResponse(req) + + default: + return mockSCIMErrorResponse(req, http.StatusMethodNotAllowed, "user method is unsupported") + } +} + +func (m *mockSCIMTransport) handleGroups(req *http.Request, segments []string, body []byte) *http.Response { + switch req.Method { + case http.MethodGet: + if len(segments) != 1 { + return mockSCIMErrorResponse(req, http.StatusMethodNotAllowed, "individual group reads are unsupported") + } + resources := make([]ScimGroup, 0, len(m.groups)) + for _, group := range m.groups { + resources = append(resources, group) + } + sort.Slice(resources, func(i, j int) bool { return resources[i].ID < resources[j].ID }) + return mockSCIMListResponse(req, resources, m.pageSize) + + case http.MethodPost: + if len(segments) != 1 { + return mockSCIMErrorResponse(req, http.StatusNotFound, "group collection path is invalid") + } + var group ScimGroup + err := json.Unmarshal(body, &group) + if err != nil { + return mockSCIMErrorResponse(req, http.StatusBadRequest, "group payload is invalid JSON") + } + violation := m.validateGroupPayload(body, group) + if violation != "" { + m.violations = append(m.violations, violation) + return mockSCIMErrorResponse(req, http.StatusBadRequest, violation) + } + status := m.failCreates[group.ExternalID] + if status != 0 { + return mockSCIMErrorResponse(req, status, "injected group creation failure") + } + m.nextID++ + group.ID = fmt.Sprintf("remote-group-%d", m.nextID) + group.Meta = newMockMeta("Group", "/Groups/"+group.ID, m.nextID) + m.groups[group.ID] = group + return mockSCIMJSONResponse(req, http.StatusCreated, group) + + case http.MethodPut: + if len(segments) != 2 { + return mockSCIMErrorResponse(req, http.StatusNotFound, "group resource path is invalid") + } + _, ok := m.groups[segments[1]] + if !ok { + return mockSCIMErrorResponse(req, http.StatusNotFound, "group does not exist") + } + var group ScimGroup + err := json.Unmarshal(body, &group) + if err != nil { + return mockSCIMErrorResponse(req, http.StatusBadRequest, "group payload is invalid JSON") + } + violation := m.validateGroupPayload(body, group) + if violation != "" { + m.violations = append(m.violations, violation) + return mockSCIMErrorResponse(req, http.StatusBadRequest, violation) + } + m.nextID++ + group.ID = segments[1] + group.Meta = newMockMeta("Group", "/Groups/"+group.ID, m.nextID) + m.groups[group.ID] = group + return mockSCIMJSONResponse(req, http.StatusOK, group) + + case http.MethodDelete: + if len(segments) != 2 { + return mockSCIMErrorResponse(req, http.StatusNotFound, "group resource path is invalid") + } + _, ok := m.groups[segments[1]] + if !ok { + return mockSCIMErrorResponse(req, http.StatusNotFound, "group does not exist") + } + delete(m.groups, segments[1]) + return mockSCIMNoContentResponse(req) + + default: + return mockSCIMErrorResponse(req, http.StatusMethodNotAllowed, "group method is unsupported") + } +} + +func validateUserPayload(body []byte, user ScimUser) string { + if violation := validateResourcePayload(body, user.ScimResourceData, scimUserSchema); violation != "" { + return violation + } + if user.UserName == "" { + return "SCIM user payload omitted userName" + } + return "" +} + +func (m *mockSCIMTransport) validateGroupPayload(body []byte, group ScimGroup) string { + if violation := validateResourcePayload(body, group.ScimResourceData, scimGroupSchema); violation != "" { + return violation + } + if group.Display == "" { + return "SCIM group payload omitted displayName" + } + for _, member := range group.Members { + if _, ok := m.users[member.Value]; !ok { + return fmt.Sprintf("SCIM group referenced unknown user %q", member.Value) + } + } + return "" +} + +func validateResourcePayload(body []byte, resource ScimResourceData, expectedSchema string) string { + if resource.ExternalID == "" { + return "SCIM resource payload omitted externalId" + } + if !slices.Contains(resource.Schemas, expectedSchema) { + return fmt.Sprintf("SCIM resource payload omitted schema %q", expectedSchema) + } + + var raw map[string]json.RawMessage + err := json.Unmarshal(body, &raw) + if err != nil { + return "SCIM resource payload is invalid JSON" + } + _, ok := raw["id"] + if ok { + return "SCIM write payload included read-only id" + } + _, ok = raw["meta"] + if ok { + return "SCIM write payload included read-only meta" + } + + return "" +} + +func mockSCIMListResponse[T any](req *http.Request, resources []T, pageSize int) *http.Response { + startIndex, err := strconv.Atoi(req.URL.Query().Get("startIndex")) + if err != nil || startIndex < 1 { + return mockSCIMErrorResponse(req, http.StatusBadRequest, "startIndex must be a one-based integer") + } + count, err := strconv.Atoi(req.URL.Query().Get("count")) + if err != nil || count < 1 { + return mockSCIMErrorResponse(req, http.StatusBadRequest, "count must be a positive integer") + } + if pageSize > 0 && pageSize < count { + count = pageSize + } + + start := min(startIndex-1, len(resources)) + end := min(start+count, len(resources)) + page := resources[start:end] + return mockSCIMJSONResponse(req, http.StatusOK, struct { + Schemas []string `json:"schemas"` + Resources []T `json:"Resources"` + TotalResults int `json:"totalResults"` + StartIndex int `json:"startIndex"` + ItemsPerPage int `json:"itemsPerPage"` + }{ + Schemas: []string{scimListResponseSchema}, + Resources: page, + TotalResults: len(resources), + StartIndex: startIndex, + ItemsPerPage: len(page), + }) +} + +func mockSCIMJSONResponse(req *http.Request, status int, payload any) *http.Response { + body, err := json.Marshal(payload) + if err != nil { + panic(err) + } + + return &http.Response{ + StatusCode: status, + Header: http.Header{"Content-Type": []string{mockSCIMRequestContentType}}, + Body: io.NopCloser(bytes.NewReader(body)), + ContentLength: int64(len(body)), + Request: req, + } +} + +func mockSCIMErrorResponse(req *http.Request, status int, detail string) *http.Response { + return mockSCIMJSONResponse(req, status, struct { + Schemas []string `json:"schemas"` + Status string `json:"status"` + Detail string `json:"detail"` + }{ + Schemas: []string{scimErrorResponseSchema}, + Status: strconv.Itoa(status), + Detail: detail, + }) +} + +func mockSCIMNoContentResponse(req *http.Request) *http.Response { + return &http.Response{ + StatusCode: http.StatusNoContent, + Header: make(http.Header), + Body: http.NoBody, + Request: req, + } +} + +func newMockMeta(resourceType, resourcePath string, version int) *ScimResourceMeta { + now := time.Now().UTC() + return &ScimResourceMeta{ + Location: mockSCIMEndpoint + resourcePath, + ResourceType: resourceType, + Created: now, + LastModified: now, + Version: fmt.Sprintf(`W/"%d"`, version), + } +} + +func (m *mockSCIMTransport) seedUser(user ScimUser) { + m.mu.Lock() + defer m.mu.Unlock() + m.users[user.ID] = user +} + +func (m *mockSCIMTransport) seedGroup(group ScimGroup) { + m.mu.Lock() + defer m.mu.Unlock() + m.groups[group.ID] = group +} + +func (m *mockSCIMTransport) user(id string) *ScimUser { + m.mu.Lock() + defer m.mu.Unlock() + + user, ok := m.users[id] + if !ok { + return nil + } + + return &user +} + +func (m *mockSCIMTransport) group(id string) *ScimGroup { + m.mu.Lock() + defer m.mu.Unlock() + + group, ok := m.groups[id] + if !ok { + return nil + } + + return &group +} + +func (m *mockSCIMTransport) usersSnapshot() map[string]ScimUser { + m.mu.Lock() + defer m.mu.Unlock() + return cloneMap(m.users) +} + +func (m *mockSCIMTransport) groupsSnapshot() map[string]ScimGroup { + m.mu.Lock() + defer m.mu.Unlock() + return cloneMap(m.groups) +} + +func (m *mockSCIMTransport) requestsSnapshot() []mockSCIMRequest { + m.mu.Lock() + defer m.mu.Unlock() + return slices.Clone(m.requests) +} + +func (m *mockSCIMTransport) requireCompliant(t *testing.T) { + t.Helper() + m.mu.Lock() + defer m.mu.Unlock() + require.Empty(t, m.violations) +} + +func cloneMap[K comparable, V any](input map[K]V) map[K]V { + result := make(map[K]V, len(input)) + maps.Copy(result, input) + return result +} + +func firstRequestIndex(requests []mockSCIMRequest, method, path string) int { + for i, request := range requests { + if request.method == method && request.path == path { + return i + } + } + return -1 +} + +func lastRequestIndex(requests []mockSCIMRequest, method, path string) int { + for i := len(requests) - 1; i >= 0; i-- { + if requests[i].method == method && requests[i].path == path { + return i + } + } + return -1 +} + +func countRequests(requests []mockSCIMRequest, method, path string) int { + count := 0 + for _, request := range requests { + if request.method == method && request.path == path { + count++ + } + } + return count +} + +func countRequestsWithPrefix(requests []mockSCIMRequest, method, pathPrefix string) int { + count := 0 + for _, request := range requests { + if request.method == method && strings.HasPrefix(request.path, pathPrefix) { + count++ + } + } + return count +} + +func countMutationRequests(requests []mockSCIMRequest) int { + count := 0 + for _, request := range requests { + if request.method == http.MethodPost || request.method == http.MethodPut || request.method == http.MethodDelete { + count++ + } + } + return count +} + +func queryValues(requests []mockSCIMRequest, method, path, key string) []string { + values := make([]string, 0) + for _, request := range requests { + if request.method == method && request.path == path { + values = append(values, request.query.Get(key)) + } + } + return values +} diff --git a/backend/internal/service/oidc_service.go b/backend/internal/service/oidc_service.go index 1a788ad5..4717136f 100644 --- a/backend/internal/service/oidc_service.go +++ b/backend/internal/service/oidc_service.go @@ -43,7 +43,7 @@ type OidcService struct { jwtService *JwtService previewBuilder oidcClientPreviewBuilder metadataRefresher metadataRefresher - scimService *ScimService + scimSyncScheduler ScimSyncScheduler httpClient *http.Client fileStorage storage.FileStorage @@ -62,7 +62,7 @@ func NewOidcService( jwtService *JwtService, previewBuilder oidcClientPreviewBuilder, metadataRefresher metadataRefresher, - scimService *ScimService, + scimSyncScheduler ScimSyncScheduler, httpClient *http.Client, fileStorage storage.FileStorage, ) (s *OidcService, err error) { @@ -71,7 +71,7 @@ func NewOidcService( jwtService: jwtService, previewBuilder: previewBuilder, metadataRefresher: metadataRefresher, - scimService: scimService, + scimSyncScheduler: scimSyncScheduler, httpClient: httpClient, fileStorage: fileStorage, } @@ -623,7 +623,9 @@ func (s *OidcService) UpdateAllowedUserGroups(ctx context.Context, id string, in return model.OidcClient{}, err } - s.scimService.ScheduleSync() + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) + } return client, nil } @@ -1024,19 +1026,3 @@ func oidcClientImagePath(clientID string, suffix string, extension string) strin } return path.Join("oidc-client-images", storageID+suffix+"."+extension) } - -func (s *OidcService) GetClientScimServiceProvider(ctx context.Context, clientID string) (model.ScimServiceProvider, error) { - var provider model.ScimServiceProvider - err := s.db. - WithContext(ctx). - First(&provider, "oidc_client_id = ?", clientID). - Error - if err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) { - return model.ScimServiceProvider{}, apperror.NotFound("SCIM service provider") - } - return model.ScimServiceProvider{}, err - } - - return provider, nil -} diff --git a/backend/internal/service/scheduler.go b/backend/internal/service/scheduler.go deleted file mode 100644 index b53f195b..00000000 --- a/backend/internal/service/scheduler.go +++ /dev/null @@ -1,25 +0,0 @@ -package service - -import ( - "context" - - backoff "github.com/cenkalti/backoff/v5" - "github.com/go-co-op/gocron/v2" -) - -// RegisterJobOpts holds optional configuration for registering a scheduled job. -type RegisterJobOpts struct { - // RunImmediately runs the job immediately after registration. - RunImmediately bool - // ExtraOptions are additional gocron job options. - ExtraOptions []gocron.JobOption - // BackOff is an optional backoff strategy. If non-nil, the job will be wrapped - // with automatic retry logic using the provided backoff on transient failures. - BackOff backoff.BackOff -} - -// Scheduler is an interface for registering and managing background jobs. -type Scheduler interface { - RegisterJob(ctx context.Context, name string, def gocron.JobDefinition, job func(ctx context.Context) error, opts RegisterJobOpts) error - RemoveJob(name string) error -} diff --git a/backend/internal/service/scim_service_test.go b/backend/internal/service/scim_service_test.go deleted file mode 100644 index 5deba203..00000000 --- a/backend/internal/service/scim_service_test.go +++ /dev/null @@ -1,61 +0,0 @@ -package service - -import ( - "testing" - - "github.com/pocket-id/pocket-id/backend/internal/apperror" - "github.com/pocket-id/pocket-id/backend/internal/dto" - "github.com/pocket-id/pocket-id/backend/internal/model" - testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing" - "github.com/stretchr/testify/require" -) - -func TestScimServiceProviderOperationsReturnSpecificNotFoundErrors(t *testing.T) { - service := NewScimService(testutils.NewDatabaseForTest(t), nil, nil) - - _, err := service.CreateServiceProvider(t.Context(), &dto.ScimServiceProviderCreateDTO{ - Endpoint: "https://scim.example.com", - OidcClientID: "missing-client", - }) - require.True(t, apperror.IsCode(err, apperror.CodeNotFound)) - - _, err = service.GetServiceProvider(t.Context(), "missing-provider") - require.True(t, apperror.IsCode(err, apperror.CodeNotFound)) - - err = service.DeleteServiceProvider(t.Context(), "missing-provider") - require.True(t, apperror.IsCode(err, apperror.CodeNotFound)) -} - -func TestScimServiceProviderCreateAndUpdate(t *testing.T) { - db := testutils.NewDatabaseForTest(t) - service := NewScimService(db, nil, nil) - - // Create two clients so provider creation and reassignment both satisfy the foreign key - require.NoError(t, db.Create(&[]model.OidcClient{ - {Base: model.Base{ID: "client-1"}, Name: "Client 1"}, - {Base: model.Base{ID: "client-2"}, Name: "Client 2"}, - }).Error) - - // Create the provider with its initial client in one transaction - provider, err := service.CreateServiceProvider(t.Context(), &dto.ScimServiceProviderCreateDTO{ - Endpoint: "https://scim.example.com/v1", - OidcClientID: "client-1", - }) - require.NoError(t, err) - require.NotEmpty(t, provider.ID) - - // Move the provider to the second client in one transaction - provider, err = service.UpdateServiceProvider(t.Context(), provider.ID, &dto.ScimServiceProviderCreateDTO{ - Endpoint: "https://scim.example.com/v2", - OidcClientID: "client-2", - }) - require.NoError(t, err) - require.Equal(t, "https://scim.example.com/v2", provider.Endpoint) - require.Equal(t, "client-2", provider.OidcClientID) - - // Verify the committed provider retains both updated values - persisted, err := service.GetServiceProvider(t.Context(), provider.ID) - require.NoError(t, err) - require.Equal(t, provider.Endpoint, persisted.Endpoint) - require.Equal(t, provider.OidcClientID, persisted.OidcClientID) -} diff --git a/backend/internal/service/scim_sync.go b/backend/internal/service/scim_sync.go new file mode 100644 index 00000000..070dba31 --- /dev/null +++ b/backend/internal/service/scim_sync.go @@ -0,0 +1,10 @@ +package service + +import ( + "context" +) + +// ScimSyncScheduler schedules a cluster-wide SCIM synchronization after application data changes +type ScimSyncScheduler interface { + ScheduleSync(ctx context.Context) +} diff --git a/backend/internal/service/user_group_service.go b/backend/internal/service/user_group_service.go index 39f850e2..dc2e7db2 100644 --- a/backend/internal/service/user_group_service.go +++ b/backend/internal/service/user_group_service.go @@ -16,12 +16,12 @@ import ( ) type UserGroupService struct { - db *gorm.DB - scimService *ScimService + db *gorm.DB + scimSyncScheduler ScimSyncScheduler } -func NewUserGroupService(db *gorm.DB, scimService *ScimService) *UserGroupService { - return &UserGroupService{db: db, scimService: scimService} +func NewUserGroupService(db *gorm.DB, scimSyncScheduler ScimSyncScheduler) *UserGroupService { + return &UserGroupService{db: db, scimSyncScheduler: scimSyncScheduler} } func (s *UserGroupService) List(ctx context.Context, name string, listRequestOptions utils.ListRequestOptions) (groups []model.UserGroup, response utils.PaginationResponse, err error) { @@ -102,15 +102,23 @@ func (s *UserGroupService) Delete(ctx context.Context, cfg *appconfig.AppConfigM return err } - if s.scimService != nil { - s.scimService.ScheduleSync() + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) } return nil } func (s *UserGroupService) Create(ctx context.Context, input dto.UserGroupCreateDto) (group model.UserGroup, err error) { - return s.CreateInternal(ctx, input, s.db) + group, err = s.CreateInternal(ctx, input, s.db) + if err != nil { + return model.UserGroup{}, err + } + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) + } + + return group, nil } // CreateInternal creates a user group within an existing transaction @@ -136,10 +144,6 @@ func (s *UserGroupService) CreateInternal(ctx context.Context, input dto.UserGro return model.UserGroup{}, err } - if s.scimService != nil { - s.scimService.ScheduleSync() - } - return group, nil } @@ -158,6 +162,9 @@ func (s *UserGroupService) Update(ctx context.Context, cfg *appconfig.AppConfigM if err != nil { return model.UserGroup{}, err } + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) + } return group, nil } @@ -196,10 +203,6 @@ func (s *UserGroupService) updateInternal(ctx context.Context, id string, input return model.UserGroup{}, err } - if s.scimService != nil { - s.scimService.ScheduleSync() - } - return group, nil } @@ -218,6 +221,9 @@ func (s *UserGroupService) UpdateUsers(ctx context.Context, id string, userIds [ if err != nil { return model.UserGroup{}, err } + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) + } return group, nil } @@ -264,10 +270,6 @@ func (s *UserGroupService) UpdateUsersInternal(ctx context.Context, id string, u return model.UserGroup{}, err } - if s.scimService != nil { - s.scimService.ScheduleSync() - } - return group, nil } @@ -347,8 +349,8 @@ func (s *UserGroupService) UpdateAllowedOidcClient(ctx context.Context, id strin return model.UserGroup{}, err } - if s.scimService != nil { - s.scimService.ScheduleSync() + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) } return group, nil diff --git a/backend/internal/service/user_service.go b/backend/internal/service/user_service.go index 4f5b7842..38d7f0fc 100644 --- a/backend/internal/service/user_service.go +++ b/backend/internal/service/user_service.go @@ -32,18 +32,18 @@ type UserService struct { auditLogService *AuditLogService customClaimService *CustomClaimService appImagesService *AppImagesService - scimService *ScimService + scimSyncScheduler ScimSyncScheduler fileStorage storage.FileStorage } -func NewUserService(db *gorm.DB, jwtService *JwtService, auditLogService *AuditLogService, customClaimService *CustomClaimService, appImagesService *AppImagesService, scimService *ScimService, fileStorage storage.FileStorage) *UserService { +func NewUserService(db *gorm.DB, jwtService *JwtService, auditLogService *AuditLogService, customClaimService *CustomClaimService, appImagesService *AppImagesService, scimSyncScheduler ScimSyncScheduler, fileStorage storage.FileStorage) *UserService { return &UserService{ db: db, jwtService: jwtService, auditLogService: auditLogService, customClaimService: customClaimService, appImagesService: appImagesService, - scimService: scimService, + scimSyncScheduler: scimSyncScheduler, fileStorage: fileStorage, } } @@ -201,6 +201,9 @@ func (s *UserService) DeleteUser(ctx context.Context, dbConfig *appconfig.AppCon if err != nil { return fmt.Errorf("failed to delete user '%s': %w", userID, err) } + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) + } // Storage operations must be executed outside of a transaction profilePicturePath := path.Join("profile-pictures", userID+".png") @@ -242,10 +245,6 @@ func (s *UserService) DeleteUserInternal(ctx context.Context, cfg *appconfig.App return fmt.Errorf("failed to delete user: %w", err) } - if s.scimService != nil { - s.scimService.ScheduleSync() - } - return nil } @@ -264,6 +263,9 @@ func (s *UserService) CreateUser(ctx context.Context, dbConfig *appconfig.AppCon if err != nil { return model.User{}, err } + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) + } return user, nil } @@ -327,7 +329,7 @@ func (s *UserService) createUserInternal(ctx context.Context, input dto.UserCrea // Bump the UpdatedAt timestamp of the groups the new user was added to // This is necessary for SCIM to work with the newly-created user, or groups may not be synced via SCIM if len(userGroups) > 0 { - err = s.touchUserGroups(ctx, tx, groupIDs(userGroups)) + err = s.touchUserGroups(ctx, tx, userGroupIDs(userGroups)) if err != nil { return model.User{}, err } @@ -348,11 +350,16 @@ func (s *UserService) createUserInternal(ctx context.Context, input dto.UserCrea } } - if s.scimService != nil { - s.scimService.ScheduleSync() + return user, nil +} + +func userGroupIDs(groups []model.UserGroup) []string { + ids := make([]string, len(groups)) + for i, group := range groups { + ids[i] = group.ID } - return user, nil + return ids } func (s *UserService) applyDefaultGroups(ctx context.Context, user *model.User, tx *gorm.DB, cfg *appconfig.AppConfigModel) error { @@ -451,6 +458,9 @@ func (s *UserService) UpdateUser(ctx context.Context, cfg *appconfig.AppConfigMo if err != nil { return model.User{}, err } + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) + } return user, nil } @@ -530,10 +540,6 @@ func (s *UserService) UpdateUserInternal(ctx context.Context, cfg *appconfig.App return user, err } - if s.scimService != nil { - s.scimService.ScheduleSync() - } - return user, nil } @@ -592,8 +598,8 @@ func (s *UserService) UpdateUserGroups(ctx context.Context, id string, userGroup return model.User{}, err } - if s.scimService != nil { - s.scimService.ScheduleSync() + if s.scimSyncScheduler != nil { + s.scimSyncScheduler.ScheduleSync(ctx) } return user, nil @@ -662,9 +668,5 @@ func (s *UserService) DisableUserInternal(ctx context.Context, tx *gorm.DB, user return err } - if s.scimService != nil { - s.scimService.ScheduleSync() - } - return nil } diff --git a/backend/internal/usersignup/module.go b/backend/internal/usersignup/module.go index d7b174f1..86db5900 100644 --- a/backend/internal/usersignup/module.go +++ b/backend/internal/usersignup/module.go @@ -27,6 +27,11 @@ type UserCreator interface { CreateUserInternal(ctx context.Context, dbConfig *appconfig.AppConfigModel, input dto.UserCreateDto, isLdapSync bool, tx *gorm.DB) (model.User, error) } +// ScimSyncScheduler schedules SCIM after the signup transaction has committed +type ScimSyncScheduler interface { + ScheduleSync(ctx context.Context) +} + type Dependencies struct { DB *gorm.DB Actors *local.Host @@ -35,6 +40,7 @@ type Dependencies struct { AuditLog AuditLogger UserCreator UserCreator AppConfig appconfig.AppConfigResolver + ScimSync ScimSyncScheduler } type Module struct { diff --git a/backend/internal/usersignup/service.go b/backend/internal/usersignup/service.go index 2b0ffc5d..b7a32bc1 100644 --- a/backend/internal/usersignup/service.go +++ b/backend/internal/usersignup/service.go @@ -31,6 +31,7 @@ type Service struct { userCreator UserCreator signer TokenService auditLog AuditLogger + scimSync ScimSyncScheduler } func newService(deps Dependencies, actorService *actor.Service) *Service { @@ -40,6 +41,7 @@ func newService(deps Dependencies, actorService *actor.Service) *Service { userCreator: deps.UserCreator, signer: deps.Signer, auditLog: deps.AuditLog, + scimSync: deps.ScimSync, } } @@ -126,6 +128,9 @@ func (s *Service) createSignedUpUser(ctx context.Context, config *appconfig.AppC if err != nil { return model.User{}, "", err } + if s.scimSync != nil { + s.scimSync.ScheduleSync(ctx) + } return user, accessToken, nil } @@ -189,6 +194,9 @@ func (s *Service) SignUpInitialAdmin(ctx context.Context, config *appconfig.AppC if err != nil { return model.User{}, "", err } + if s.scimSync != nil { + s.scimSync.ScheduleSync(ctx) + } return user, token, nil }