mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-05 17:29:04 +02:00
feat: add OIDC back-channel logout (#1734)
Co-authored-by: Alessandro (Ale) Segala <43508+ItalyPaleAle@users.noreply.github.com> Co-authored-by: Elias Schneider <login@eliasschneider.com>
This commit is contained in:
co-authored by
Alessandro Segala
Elias Schneider
parent
23c4825abd
commit
2075de3234
@@ -41,16 +41,27 @@ type ScimSyncScheduler interface {
|
||||
ScheduleSync(ctx context.Context)
|
||||
}
|
||||
|
||||
// BackchannelLogoutNotifier tells OIDC clients to end their sessions for users the sync deprovisions
|
||||
type BackchannelLogoutNotifier interface {
|
||||
// PrepareUserNotifications resolves the notifications within the sync transaction and returns a function that delivers them, which must only be called after the transaction has committed
|
||||
PrepareUserNotifications(ctx context.Context, tx *gorm.DB, userIDs []string) (func(), error)
|
||||
|
||||
// NotifyLostGroupAccess delivers logout tokens to group-restricted clients the given users can no longer access, and must be called after the transaction has committed
|
||||
// The sync passes an empty client ID, matching on the users alone
|
||||
NotifyLostGroupAccess(ctx context.Context, userIDs []string, clientID string)
|
||||
}
|
||||
|
||||
type Dependencies struct {
|
||||
DB *gorm.DB
|
||||
Actors francishost.Host
|
||||
HTTPClient *http.Client
|
||||
FileStorage storage.FileStorage
|
||||
|
||||
Users UserSyncer
|
||||
Groups GroupSyncer
|
||||
AppConfig appconfig.AppConfigResolver
|
||||
ScimSync ScimSyncScheduler
|
||||
Users UserSyncer
|
||||
Groups GroupSyncer
|
||||
AppConfig appconfig.AppConfigResolver
|
||||
ScimSync ScimSyncScheduler
|
||||
BackchannelLogout BackchannelLogoutNotifier
|
||||
|
||||
// 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
|
||||
|
||||
@@ -8,9 +8,11 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"maps"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"path"
|
||||
"slices"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
@@ -31,13 +33,14 @@ import (
|
||||
// Service performs the actual LDAP synchronization
|
||||
// It is deliberately free of any actor concern: the sync actor only decides when a sync runs, while the reconciliation logic lives here and is called directly by the manual "sync now" endpoint too
|
||||
type Service struct {
|
||||
db *gorm.DB
|
||||
httpClient *http.Client
|
||||
users UserSyncer
|
||||
groups GroupSyncer
|
||||
scimSync ScimSyncScheduler
|
||||
fileStorage storage.FileStorage
|
||||
clientFactory func(dbConfig *appconfig.AppConfigModel) (ldapClient, error)
|
||||
db *gorm.DB
|
||||
httpClient *http.Client
|
||||
users UserSyncer
|
||||
groups GroupSyncer
|
||||
scimSync ScimSyncScheduler
|
||||
backchannelLogout BackchannelLogoutNotifier
|
||||
fileStorage storage.FileStorage
|
||||
clientFactory func(dbConfig *appconfig.AppConfigModel) (ldapClient, error)
|
||||
}
|
||||
|
||||
type savePicture struct {
|
||||
@@ -73,12 +76,13 @@ type ldapClient interface {
|
||||
|
||||
func newService(deps Dependencies) *Service {
|
||||
service := &Service{
|
||||
db: deps.DB,
|
||||
httpClient: deps.HTTPClient,
|
||||
users: deps.Users,
|
||||
groups: deps.Groups,
|
||||
scimSync: deps.ScimSync,
|
||||
fileStorage: deps.FileStorage,
|
||||
db: deps.DB,
|
||||
httpClient: deps.HTTPClient,
|
||||
users: deps.Users,
|
||||
groups: deps.Groups,
|
||||
scimSync: deps.ScimSync,
|
||||
backchannelLogout: deps.BackchannelLogout,
|
||||
fileStorage: deps.FileStorage,
|
||||
}
|
||||
|
||||
service.clientFactory = service.createClient
|
||||
@@ -129,13 +133,13 @@ func (s *Service) SyncAll(ctx context.Context, dbConfig *appconfig.AppConfigMode
|
||||
defer tx.Rollback()
|
||||
|
||||
// Reconcile users
|
||||
savePictures, deleteFiles, err := s.reconcileUsers(ctx, tx, desiredState.users, desiredState.userIDs, dbConfig)
|
||||
savePictures, deleteFiles, notifyLogout, err := s.reconcileUsers(ctx, tx, desiredState.users, desiredState.userIDs, dbConfig)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to sync users: %w", err)
|
||||
}
|
||||
|
||||
// Reconcile groups
|
||||
err = s.reconcileGroups(ctx, tx, desiredState.groups, desiredState.groupIDs, dbConfig)
|
||||
usersRemovedFromGroups, err := s.reconcileGroups(ctx, tx, desiredState.groups, desiredState.groupIDs, dbConfig)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to sync groups: %w", err)
|
||||
}
|
||||
@@ -151,6 +155,12 @@ func (s *Service) SyncAll(ctx context.Context, dbConfig *appconfig.AppConfigMode
|
||||
s.scimSync.ScheduleSync(ctx)
|
||||
}
|
||||
|
||||
// Tell OIDC clients to end the sessions of users the sync deprovisioned or removed from a group, now that the transaction has committed
|
||||
notifyLogout()
|
||||
if s.backchannelLogout != nil {
|
||||
s.backchannelLogout.NotifyLostGroupAccess(ctx, usersRemovedFromGroups, "")
|
||||
}
|
||||
|
||||
// Now that we've committed the transaction, we can perform operations on the storage layer
|
||||
// First, save all new pictures
|
||||
for _, sp := range savePictures {
|
||||
@@ -422,18 +432,30 @@ func (s *Service) resolveGroupMemberUsername(ctx context.Context, client ldapCli
|
||||
return norm.NFC.String(username)
|
||||
}
|
||||
|
||||
func (s *Service) reconcileGroups(ctx context.Context, tx *gorm.DB, desiredGroups []ldapDesiredGroup, ldapGroupIDs map[string]struct{}, dbConfig *appconfig.AppConfigModel) error {
|
||||
// reconcileGroups returns the IDs of the users this sync removed from a group, which may cost them access to clients restricted to it
|
||||
func (s *Service) reconcileGroups(ctx context.Context, tx *gorm.DB, desiredGroups []ldapDesiredGroup, ldapGroupIDs map[string]struct{}, dbConfig *appconfig.AppConfigModel) ([]string, error) {
|
||||
// Load the current LDAP-managed state from the database
|
||||
ldapGroupsInDB, ldapGroupsByID, err := s.loadLDAPGroupsInDB(ctx, tx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to fetch groups from database: %w", err)
|
||||
return nil, fmt.Errorf("failed to fetch groups from database: %w", err)
|
||||
}
|
||||
|
||||
_, _, ldapUsersByUsername, err := s.loadLDAPUsersInDB(ctx, tx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to fetch users from database: %w", err)
|
||||
return nil, fmt.Errorf("failed to fetch users from database: %w", err)
|
||||
}
|
||||
|
||||
// Capture the memberships before they are reconciled, as removals are only visible by comparing against the desired state
|
||||
var membersByGroup map[string][]string
|
||||
if s.backchannelLogout != nil {
|
||||
membersByGroup, err = s.loadGroupMembers(ctx, tx, ldapGroupsInDB)
|
||||
if err != nil {
|
||||
// Notifications are best effort and must never fail the sync
|
||||
slog.Warn("Failed to load group members to notify for back-channel logout", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
removedMembers := map[string]struct{}{}
|
||||
|
||||
// Apply creates and updates to match the desired LDAP group state
|
||||
for _, desiredGroup := range desiredGroups {
|
||||
memberUserIDs := make([]string, 0, len(desiredGroup.memberUsernames))
|
||||
@@ -451,26 +473,28 @@ func (s *Service) reconcileGroups(ctx context.Context, tx *gorm.DB, desiredGroup
|
||||
if databaseGroup.ID == "" {
|
||||
newGroup, err := s.groups.CreateInternal(ctx, desiredGroup.input, tx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create group '%s': %w", desiredGroup.input.Name, err)
|
||||
return nil, fmt.Errorf("failed to create group '%s': %w", desiredGroup.input.Name, err)
|
||||
}
|
||||
ldapGroupsByID[desiredGroup.ldapID] = newGroup
|
||||
|
||||
_, err = s.groups.UpdateUsersInternal(ctx, newGroup.ID, memberUserIDs, tx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to sync users for group '%s': %w", desiredGroup.input.Name, err)
|
||||
return nil, fmt.Errorf("failed to sync users for group '%s': %w", desiredGroup.input.Name, err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
_, err = s.groups.UpdateInternal(ctx, dbConfig, databaseGroup.ID, desiredGroup.input, true, tx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update group '%s': %w", desiredGroup.input.Name, err)
|
||||
return nil, fmt.Errorf("failed to update group '%s': %w", desiredGroup.input.Name, err)
|
||||
}
|
||||
|
||||
_, err = s.groups.UpdateUsersInternal(ctx, databaseGroup.ID, memberUserIDs, tx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to sync users for group '%s': %w", desiredGroup.input.Name, err)
|
||||
return nil, fmt.Errorf("failed to sync users for group '%s': %w", desiredGroup.input.Name, err)
|
||||
}
|
||||
|
||||
collectRemovedMembers(removedMembers, membersByGroup[databaseGroup.ID], memberUserIDs)
|
||||
}
|
||||
|
||||
// Delete groups that are no longer present in LDAP
|
||||
@@ -488,21 +512,68 @@ func (s *Service) reconcileGroups(ctx context.Context, tx *gorm.DB, desiredGroup
|
||||
Delete(&model.UserGroup{}, "ldap_id = ?", *group.LdapID).
|
||||
Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete group '%s': %w", group.Name, err)
|
||||
return nil, fmt.Errorf("failed to delete group '%s': %w", group.Name, err)
|
||||
}
|
||||
|
||||
slog.Info("Deleted group", slog.String("group", group.Name))
|
||||
|
||||
// Deleting the group removes every member from it
|
||||
collectRemovedMembers(removedMembers, membersByGroup[group.ID], nil)
|
||||
}
|
||||
|
||||
return nil
|
||||
return slices.Collect(maps.Keys(removedMembers)), nil
|
||||
}
|
||||
|
||||
// collectRemovedMembers adds the previous members that are not part of the group's new member list to removedMembers
|
||||
func collectRemovedMembers(removedMembers map[string]struct{}, previousMemberIDs []string, memberIDs []string) {
|
||||
for _, previousMemberID := range previousMemberIDs {
|
||||
if !slices.Contains(memberIDs, previousMemberID) {
|
||||
removedMembers[previousMemberID] = struct{}{}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// loadGroupMembers returns the IDs of the users that are currently members of the given groups, indexed by group ID
|
||||
func (s *Service) loadGroupMembers(ctx context.Context, tx *gorm.DB, groups []model.UserGroup) (map[string][]string, error) {
|
||||
if len(groups) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
groupIDs := make([]string, len(groups))
|
||||
for i, group := range groups {
|
||||
groupIDs[i] = group.ID
|
||||
}
|
||||
|
||||
var memberships []struct {
|
||||
UserGroupID string
|
||||
UserID string
|
||||
}
|
||||
err := tx.
|
||||
WithContext(ctx).
|
||||
Table("user_groups_users").
|
||||
Select("user_group_id", "user_id").
|
||||
Where("user_group_id IN (?)", groupIDs).
|
||||
Find(&memberships).
|
||||
Error
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
membersByGroup := make(map[string][]string, len(groups))
|
||||
for _, membership := range memberships {
|
||||
membersByGroup[membership.UserGroupID] = append(membersByGroup[membership.UserGroupID], membership.UserID)
|
||||
}
|
||||
return membersByGroup, nil
|
||||
}
|
||||
|
||||
//nolint:gocognit
|
||||
func (s *Service) reconcileUsers(ctx context.Context, tx *gorm.DB, desiredUsers []ldapDesiredUser, ldapUserIDs map[string]struct{}, dbConfig *appconfig.AppConfigModel) (savePictures []savePicture, deleteFiles []string, err error) {
|
||||
func (s *Service) reconcileUsers(ctx context.Context, tx *gorm.DB, desiredUsers []ldapDesiredUser, ldapUserIDs map[string]struct{}, dbConfig *appconfig.AppConfigModel) (savePictures []savePicture, deleteFiles []string, notifyLogout func(), err error) {
|
||||
notifyLogout = func() {}
|
||||
|
||||
// Load the current LDAP-managed state from the database
|
||||
ldapUsersInDB, ldapUsersByID, _, err := s.loadLDAPUsersInDB(ctx, tx)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to fetch users from database: %w", err)
|
||||
return nil, nil, nil, fmt.Errorf("failed to fetch users from database: %w", err)
|
||||
}
|
||||
|
||||
// Apply creates and updates to match the desired LDAP user state
|
||||
@@ -520,7 +591,7 @@ func (s *Service) reconcileUsers(ctx context.Context, tx *gorm.DB, desiredUsers
|
||||
Update("disabled", false).
|
||||
Error
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to enable user %s: %w", databaseUser.Username, err)
|
||||
return nil, nil, nil, fmt.Errorf("failed to enable user %s: %w", databaseUser.Username, err)
|
||||
}
|
||||
|
||||
databaseUser.Disabled = false
|
||||
@@ -534,7 +605,7 @@ func (s *Service) reconcileUsers(ctx context.Context, tx *gorm.DB, desiredUsers
|
||||
slog.Warn("Skipping creating LDAP user", slog.String("username", desiredUser.input.Username), slog.Any("error", err))
|
||||
continue
|
||||
} else if err != nil {
|
||||
return nil, nil, fmt.Errorf("error creating user '%s': %w", desiredUser.input.Username, err)
|
||||
return nil, nil, nil, fmt.Errorf("error creating user '%s': %w", desiredUser.input.Username, err)
|
||||
}
|
||||
|
||||
userID = createdUser.ID
|
||||
@@ -545,7 +616,7 @@ func (s *Service) reconcileUsers(ctx context.Context, tx *gorm.DB, desiredUsers
|
||||
slog.Warn("Skipping updating LDAP user", slog.String("username", desiredUser.input.Username), slog.Any("error", err))
|
||||
continue
|
||||
} else if err != nil {
|
||||
return nil, nil, fmt.Errorf("error updating user '%s': %w", desiredUser.input.Username, err)
|
||||
return nil, nil, nil, fmt.Errorf("error updating user '%s': %w", desiredUser.input.Username, err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -558,21 +629,36 @@ func (s *Service) reconcileUsers(ctx context.Context, tx *gorm.DB, desiredUsers
|
||||
}
|
||||
}
|
||||
|
||||
// The authorizations of deleted users are gone once the transaction commits, so the clients to notify must be resolved before deprovisioning
|
||||
// Users that a previous sync already disabled are re-disabled on every run and must not be notified again
|
||||
if s.backchannelLogout != nil {
|
||||
deprovisionedUserIDs := make([]string, 0, len(ldapUsersInDB))
|
||||
for _, user := range ldapUsersInDB {
|
||||
if !isDeprovisioned(user, ldapUserIDs) || (dbConfig.LdapSoftDeleteUsers.IsTrue() && user.Disabled) {
|
||||
continue
|
||||
}
|
||||
|
||||
deprovisionedUserIDs = append(deprovisionedUserIDs, user.ID)
|
||||
}
|
||||
|
||||
notifyLogout, err = s.backchannelLogout.PrepareUserNotifications(ctx, tx, deprovisionedUserIDs)
|
||||
if err != nil {
|
||||
// Notifications are best effort and must never fail the sync
|
||||
slog.Warn("Failed to prepare back-channel logout notifications for deprovisioned LDAP users", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
|
||||
// Disable or delete users that are no longer present in LDAP
|
||||
deleteFiles = make([]string, 0, len(ldapUsersInDB))
|
||||
for _, user := range ldapUsersInDB {
|
||||
if user.LdapID == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
if _, exists := ldapUserIDs[*user.LdapID]; exists {
|
||||
if !isDeprovisioned(user, ldapUserIDs) {
|
||||
continue
|
||||
}
|
||||
|
||||
if dbConfig.LdapSoftDeleteUsers.IsTrue() {
|
||||
err = s.users.DisableUserInternal(ctx, tx, user.ID)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("failed to disable user %s: %w", user.Username, err)
|
||||
return nil, nil, nil, fmt.Errorf("failed to disable user %s: %w", user.Username, err)
|
||||
}
|
||||
|
||||
slog.Info("Disabled user", slog.String("username", user.Username))
|
||||
@@ -582,16 +668,26 @@ func (s *Service) reconcileUsers(ctx context.Context, tx *gorm.DB, desiredUsers
|
||||
err = s.users.DeleteUserInternal(ctx, dbConfig, tx, user.ID, true)
|
||||
if err != nil {
|
||||
if apperror.IsCode(err, apperror.CodeLdapUserUpdate) {
|
||||
return nil, nil, fmt.Errorf("failed to delete user %s: LDAP user must be disabled before deletion", user.Username)
|
||||
return nil, nil, nil, fmt.Errorf("failed to delete user %s: LDAP user must be disabled before deletion", user.Username)
|
||||
}
|
||||
return nil, nil, fmt.Errorf("failed to delete user %s: %w", user.Username, err)
|
||||
return nil, nil, nil, fmt.Errorf("failed to delete user %s: %w", user.Username, err)
|
||||
}
|
||||
|
||||
slog.Info("Deleted user", slog.String("username", user.Username))
|
||||
deleteFiles = append(deleteFiles, path.Join("profile-pictures", user.ID+".png"))
|
||||
}
|
||||
|
||||
return savePictures, deleteFiles, nil
|
||||
return savePictures, deleteFiles, notifyLogout, nil
|
||||
}
|
||||
|
||||
// isDeprovisioned reports whether an LDAP-managed user is no longer present in the directory and is therefore disabled or deleted by this sync
|
||||
func isDeprovisioned(user model.User, ldapUserIDs map[string]struct{}) bool {
|
||||
if user.LdapID == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
_, exists := ldapUserIDs[*user.LdapID]
|
||||
return !exists
|
||||
}
|
||||
|
||||
func (s *Service) loadLDAPUsersInDB(ctx context.Context, tx *gorm.DB) (users []model.User, byLdapID map[string]model.User, byUsername map[string]model.User, err error) {
|
||||
|
||||
@@ -317,6 +317,39 @@ func TestLdapServiceSyncAllSetsAdminFromGroupMembership(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadGroupMembers(t *testing.T) {
|
||||
svc, db := newTestLdapService(t, nil)
|
||||
|
||||
require.NoError(t, db.Create(&model.User{Base: model.Base{ID: "user-1"}, Username: "user1"}).Error)
|
||||
require.NoError(t, db.Create(&model.User{Base: model.Base{ID: "user-2"}, Username: "user2"}).Error)
|
||||
|
||||
group := model.UserGroup{Base: model.Base{ID: "group-1"}, Name: "group1", FriendlyName: "Group 1"}
|
||||
require.NoError(t, db.Create(&group).Error)
|
||||
require.NoError(t, db.Model(&group).Association("Users").Append([]model.User{
|
||||
{Base: model.Base{ID: "user-1"}},
|
||||
{Base: model.Base{ID: "user-2"}},
|
||||
}))
|
||||
|
||||
emptyGroup := model.UserGroup{Base: model.Base{ID: "group-2"}, Name: "group2", FriendlyName: "Group 2"}
|
||||
require.NoError(t, db.Create(&emptyGroup).Error)
|
||||
|
||||
membersByGroup, err := svc.loadGroupMembers(t.Context(), db, []model.UserGroup{group, emptyGroup})
|
||||
require.NoError(t, err)
|
||||
assert.ElementsMatch(t, []string{"user-1", "user-2"}, membersByGroup["group-1"])
|
||||
assert.Empty(t, membersByGroup["group-2"])
|
||||
}
|
||||
|
||||
func TestCollectRemovedMembers(t *testing.T) {
|
||||
removedMembers := map[string]struct{}{}
|
||||
|
||||
collectRemovedMembers(removedMembers, []string{"user-1", "user-2"}, []string{"user-2", "user-3"})
|
||||
assert.Equal(t, map[string]struct{}{"user-1": {}}, removedMembers)
|
||||
|
||||
// Deleting a group passes no remaining members, and a user removed from two groups is only collected once
|
||||
collectRemovedMembers(removedMembers, []string{"user-1", "user-4"}, nil)
|
||||
assert.Equal(t, map[string]struct{}{"user-1": {}, "user-4": {}}, removedMembers)
|
||||
}
|
||||
|
||||
func newTestLdapService(t *testing.T, client ldapClient) (*Service, *gorm.DB) {
|
||||
t.Helper()
|
||||
|
||||
@@ -326,7 +359,7 @@ func newTestLdapService(t *testing.T, client ldapClient) (*Service, *gorm.DB) {
|
||||
require.NoError(t, err)
|
||||
|
||||
// The sync is exercised against the real user and group services, so the assertions below can check what actually lands in the database
|
||||
groupService := service.NewUserGroupService(db, nil)
|
||||
groupService := service.NewUserGroupService(db, nil, nil)
|
||||
userService := service.NewUserService(
|
||||
db,
|
||||
nil,
|
||||
@@ -334,6 +367,7 @@ func newTestLdapService(t *testing.T, client ldapClient) (*Service, *gorm.DB) {
|
||||
service.NewCustomClaimService(db),
|
||||
service.NewAppImagesService(map[string]string{}, fileStorage),
|
||||
nil,
|
||||
nil,
|
||||
fileStorage,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user