refactor: migrate SCIM sync to actor + add tests (#1680)

Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
This commit is contained in:
Alessandro (Ale) Segala
2026-08-18 12:54:06 -07:00
committed by GitHub
parent a099d9457c
commit 8c095afd78
33 changed files with 1919 additions and 782 deletions

View File

@@ -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

View File

@@ -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=

View File

@@ -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",

View File

@@ -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

View File

@@ -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),

View File

@@ -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
}

View File

@@ -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)

View File

@@ -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)
}

View File

@@ -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)

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}
}

View File

@@ -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)
}

View File

@@ -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

View File

@@ -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 {

View File

@@ -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;"`
}

View File

@@ -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)))
}

View File

@@ -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
}

View File

@@ -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 {

View File

@@ -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
}

View File

@@ -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"
}

View File

@@ -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))
}
}

View File

@@ -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
}

View File

@@ -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")
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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)
}

View File

@@ -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)
}

View File

@@ -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

View File

@@ -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
}

View File

@@ -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 {

View File

@@ -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
}