mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-09-08 04:01:26 +02:00
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:
committed by
GitHub
parent
a099d9457c
commit
8c095afd78
@@ -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
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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;"`
|
||||
}
|
||||
148
backend/internal/scimsync/actor.go
Normal file
148
backend/internal/scimsync/actor.go
Normal 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)))
|
||||
}
|
||||
180
backend/internal/scimsync/actor_test.go
Normal file
180
backend/internal/scimsync/actor_test.go
Normal 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
|
||||
}
|
||||
@@ -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 {
|
||||
135
backend/internal/scimsync/handler.go
Normal file
135
backend/internal/scimsync/handler.go
Normal 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
|
||||
}
|
||||
21
backend/internal/scimsync/model.go
Normal file
21
backend/internal/scimsync/model.go
Normal 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"
|
||||
}
|
||||
70
backend/internal/scimsync/module.go
Normal file
70
backend/internal/scimsync/module.go
Normal 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))
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
136
backend/internal/scimsync/service_test.go
Normal file
136
backend/internal/scimsync/service_test.go
Normal 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")
|
||||
}
|
||||
837
backend/internal/scimsync/sync_test.go
Normal file
837
backend/internal/scimsync/sync_test.go
Normal 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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
10
backend/internal/service/scim_sync.go
Normal file
10
backend/internal/service/scim_sync.go
Normal 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)
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user