diff --git a/pkg/iam/scim/bridge_runner_sync.go b/pkg/iam/scim/bridge_runner_sync.go index 6196298e6..c42dcb64b 100644 --- a/pkg/iam/scim/bridge_runner_sync.go +++ b/pkg/iam/scim/bridge_runner_sync.go @@ -32,70 +32,56 @@ import ( func (r *BridgeRunner) executeSync( ctx context.Context, - bridge *coredata.SCIMBridge, + scimBridge *coredata.SCIMBridge, scope coredata.Scoper, logger *log.Logger, -) (stats SyncStats, duration time.Duration, connector *coredata.Connector, err error) { +) (stats SyncStats, duration time.Duration, dbConnector *coredata.Connector, err error) { start := time.Now() + var ( + idp provider.Provider + token string + ) + err = r.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) error { - var syncErr error - stats, connector, syncErr = r.doSync(ctx, tx, bridge, scope, logger) - return syncErr + var err error + idp, token, dbConnector, err = r.prepareSync( + ctx, + tx, + scimBridge, + scope, + logger, + ) + + if err != nil { + return fmt.Errorf("cannot prepare sync: %w", err) + } + + return nil }, ) - - duration = time.Since(start) - return stats, duration, connector, err -} - -func (r *BridgeRunner) doSync( - ctx context.Context, - tx pg.Tx, - scimBridge *coredata.SCIMBridge, - scope coredata.Scoper, - logger *log.Logger, -) (SyncStats, *coredata.Connector, error) { - if scimBridge.ConnectorID == nil { - return SyncStats{}, nil, fmt.Errorf("bridge has no connector configured") - } - - dbConnector := &coredata.Connector{} - if err := dbConnector.LoadByID(ctx, tx, scope, *scimBridge.ConnectorID, r.encryptionKey); err != nil { - return SyncStats{}, nil, fmt.Errorf("cannot load connector: %w", err) - } - - idp, err := r.createProvider(ctx, logger, scimBridge.Type, dbConnector, scimBridge.ExcludedUserNames) if err != nil { - return SyncStats{}, nil, fmt.Errorf("cannot create provider: %w", err) - } - - var scimConfig coredata.SCIMConfiguration - if err := scimConfig.LoadByID(ctx, tx, scope, scimBridge.ScimConfigurationID); err != nil { - return SyncStats{}, nil, fmt.Errorf("cannot load SCIM configuration: %w", err) - } - - token, err := GenerateToken() - if err != nil { - return SyncStats{}, nil, fmt.Errorf("cannot generate SCIM token: %w", err) - } - - scimConfig.HashedToken = HashToken(token) - scimConfig.UpdatedAt = time.Now() - if err := scimConfig.Update(ctx, tx, scope); err != nil { - return SyncStats{}, nil, fmt.Errorf("cannot update SCIM configuration token: %w", err) + duration = time.Since(start) + return SyncStats{}, duration, nil, err } scimClient := r.createSCIMClient(logger, token) - syncer := bridge.NewBridge(idp, scimClient, bridge.WithExcludedUserNames(scimBridge.ExcludedUserNames)) - created, updated, deleted, deactivated, skipped, err := syncer.Run(ctx) - if err != nil { - return SyncStats{}, nil, fmt.Errorf("sync failed: %w", err) + syncer := bridge.NewBridge( + idp, + scimClient, + bridge.WithExcludedUserNames(scimBridge.ExcludedUserNames), + ) + + created, updated, deleted, deactivated, skipped, syncErr := syncer.Run(ctx) + duration = time.Since(start) + + if syncErr != nil { + return SyncStats{}, duration, dbConnector, fmt.Errorf("sync failed: %w", syncErr) } - stats := SyncStats{ + stats = SyncStats{ Created: created, Updated: updated, Deleted: deleted, @@ -103,7 +89,47 @@ func (r *BridgeRunner) doSync( Skipped: skipped, } - return stats, dbConnector, nil + return stats, duration, dbConnector, nil +} + +func (r *BridgeRunner) prepareSync( + ctx context.Context, + tx pg.Tx, + scimBridge *coredata.SCIMBridge, + scope coredata.Scoper, + logger *log.Logger, +) (provider.Provider, string, *coredata.Connector, error) { + if scimBridge.ConnectorID == nil { + return nil, "", nil, fmt.Errorf("bridge has no connector configured") + } + + dbConnector := &coredata.Connector{} + if err := dbConnector.LoadByID(ctx, tx, scope, *scimBridge.ConnectorID, r.encryptionKey); err != nil { + return nil, "", nil, fmt.Errorf("cannot load connector: %w", err) + } + + idp, err := r.createProvider(ctx, logger, scimBridge.Type, dbConnector, scimBridge.ExcludedUserNames) + if err != nil { + return nil, "", nil, fmt.Errorf("cannot create provider: %w", err) + } + + var scimConfig coredata.SCIMConfiguration + if err := scimConfig.LoadByID(ctx, tx, scope, scimBridge.ScimConfigurationID); err != nil { + return nil, "", nil, fmt.Errorf("cannot load SCIM configuration: %w", err) + } + + token, err := GenerateToken() + if err != nil { + return nil, "", nil, fmt.Errorf("cannot generate SCIM token: %w", err) + } + + scimConfig.HashedToken = HashToken(token) + scimConfig.UpdatedAt = time.Now() + if err := scimConfig.Update(ctx, tx, scope); err != nil { + return nil, "", nil, fmt.Errorf("cannot update SCIM configuration token: %w", err) + } + + return idp, token, dbConnector, nil } func (r *BridgeRunner) createSCIMClient(logger *log.Logger, token string) *scimclient.Client {