mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
chore: support multi-org group sync with runtime configuration (#14578)
- Implement multi-org group sync - Implement runtime configuration to change sync behavior - Legacy group sync migrated to new package
This commit is contained in:
@@ -0,0 +1,416 @@
|
||||
package idpsync
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"regexp"
|
||||
|
||||
"github.com/golang-jwt/jwt/v4"
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/coderd/runtimeconfig"
|
||||
"github.com/coder/coder/v2/coderd/util/slice"
|
||||
)
|
||||
|
||||
type GroupParams struct {
|
||||
// SyncEnabled if false will skip syncing the user's groups
|
||||
SyncEnabled bool
|
||||
MergedClaims jwt.MapClaims
|
||||
}
|
||||
|
||||
func (AGPLIDPSync) GroupSyncEnabled() bool {
|
||||
// AGPL does not support syncing groups.
|
||||
return false
|
||||
}
|
||||
|
||||
func (s AGPLIDPSync) GroupSyncSettings() runtimeconfig.RuntimeEntry[*GroupSyncSettings] {
|
||||
return s.Group
|
||||
}
|
||||
|
||||
func (s AGPLIDPSync) ParseGroupClaims(_ context.Context, _ jwt.MapClaims) (GroupParams, *HTTPError) {
|
||||
return GroupParams{
|
||||
SyncEnabled: s.GroupSyncEnabled(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s AGPLIDPSync) SyncGroups(ctx context.Context, db database.Store, user database.User, params GroupParams) error {
|
||||
// Nothing happens if sync is not enabled
|
||||
if !params.SyncEnabled {
|
||||
return nil
|
||||
}
|
||||
|
||||
// nolint:gocritic // all syncing is done as a system user
|
||||
ctx = dbauthz.AsSystemRestricted(ctx)
|
||||
|
||||
// Only care about the default org for deployment settings if the
|
||||
// legacy deployment settings exist.
|
||||
defaultOrgID := uuid.Nil
|
||||
// Default organization is configured via legacy deployment values
|
||||
if s.DeploymentSyncSettings.Legacy.GroupField != "" {
|
||||
defaultOrganization, err := db.GetDefaultOrganization(ctx)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get default organization: %w", err)
|
||||
}
|
||||
defaultOrgID = defaultOrganization.ID
|
||||
}
|
||||
|
||||
err := db.InTx(func(tx database.Store) error {
|
||||
userGroups, err := tx.GetGroups(ctx, database.GetGroupsParams{
|
||||
HasMemberID: user.ID,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get user groups: %w", err)
|
||||
}
|
||||
|
||||
// Figure out which organizations the user is a member of.
|
||||
// The "Everyone" group is always included, so we can infer organization
|
||||
// membership via the groups the user is in.
|
||||
userOrgs := make(map[uuid.UUID][]database.GetGroupsRow)
|
||||
for _, g := range userGroups {
|
||||
g := g
|
||||
userOrgs[g.Group.OrganizationID] = append(userOrgs[g.Group.OrganizationID], g)
|
||||
}
|
||||
|
||||
// For each org, we need to fetch the sync settings
|
||||
// This loop also handles any legacy settings for the default
|
||||
// organization.
|
||||
orgSettings := make(map[uuid.UUID]GroupSyncSettings)
|
||||
for orgID := range userOrgs {
|
||||
orgResolver := s.Manager.OrganizationResolver(tx, orgID)
|
||||
settings, err := s.SyncSettings.Group.Resolve(ctx, orgResolver)
|
||||
if err != nil {
|
||||
if !xerrors.Is(err, runtimeconfig.ErrEntryNotFound) {
|
||||
return xerrors.Errorf("resolve group sync settings: %w", err)
|
||||
}
|
||||
// Default to not being configured
|
||||
settings = &GroupSyncSettings{}
|
||||
}
|
||||
|
||||
// Legacy deployment settings will override empty settings.
|
||||
if orgID == defaultOrgID && settings.Field == "" {
|
||||
settings = &GroupSyncSettings{
|
||||
Field: s.Legacy.GroupField,
|
||||
LegacyNameMapping: s.Legacy.GroupMapping,
|
||||
RegexFilter: s.Legacy.GroupFilter,
|
||||
AutoCreateMissing: s.Legacy.CreateMissingGroups,
|
||||
}
|
||||
}
|
||||
orgSettings[orgID] = *settings
|
||||
}
|
||||
|
||||
// groupIDsToAdd & groupIDsToRemove are the final group differences
|
||||
// needed to be applied to user. The loop below will iterate over all
|
||||
// organizations the user is in, and determine the diffs.
|
||||
// The diffs are applied as a batch sql query, rather than each
|
||||
// organization having to execute a query.
|
||||
groupIDsToAdd := make([]uuid.UUID, 0)
|
||||
groupIDsToRemove := make([]uuid.UUID, 0)
|
||||
// For each org, determine which groups the user should land in
|
||||
for orgID, settings := range orgSettings {
|
||||
if settings.Field == "" {
|
||||
// No group sync enabled for this org, so do nothing.
|
||||
// The user can remain in their groups for this org.
|
||||
continue
|
||||
}
|
||||
|
||||
// expectedGroups is the set of groups the IDP expects the
|
||||
// user to be a member of.
|
||||
expectedGroups, err := settings.ParseClaims(orgID, params.MergedClaims)
|
||||
if err != nil {
|
||||
s.Logger.Debug(ctx, "failed to parse claims for groups",
|
||||
slog.F("organization_field", s.GroupField),
|
||||
slog.F("organization_id", orgID),
|
||||
slog.Error(err),
|
||||
)
|
||||
// Unsure where to raise this error on the UI or database.
|
||||
// TODO: This error prevents group sync, but we have no way
|
||||
// to raise this to an org admin. Come up with a solution to
|
||||
// notify the admin and user of this issue.
|
||||
continue
|
||||
}
|
||||
// Everyone group is always implied, so include it.
|
||||
expectedGroups = append(expectedGroups, ExpectedGroup{
|
||||
OrganizationID: orgID,
|
||||
GroupID: &orgID,
|
||||
})
|
||||
|
||||
// Now we know what groups the user should be in for a given org,
|
||||
// determine if we have to do any group updates to sync the user's
|
||||
// state.
|
||||
existingGroups := userOrgs[orgID]
|
||||
existingGroupsTyped := db2sdk.List(existingGroups, func(f database.GetGroupsRow) ExpectedGroup {
|
||||
return ExpectedGroup{
|
||||
OrganizationID: orgID,
|
||||
GroupID: &f.Group.ID,
|
||||
GroupName: &f.Group.Name,
|
||||
}
|
||||
})
|
||||
|
||||
add, remove := slice.SymmetricDifferenceFunc(existingGroupsTyped, expectedGroups, func(a, b ExpectedGroup) bool {
|
||||
return a.Equal(b)
|
||||
})
|
||||
|
||||
for _, r := range remove {
|
||||
if r.GroupID == nil {
|
||||
// This should never happen. All group removals come from the
|
||||
// existing set, which come from the db. All groups from the
|
||||
// database have IDs. This code is purely defensive.
|
||||
detail := "user:" + user.Username
|
||||
if r.GroupName != nil {
|
||||
detail += fmt.Sprintf(" from group %s", *r.GroupName)
|
||||
}
|
||||
return xerrors.Errorf("removal group has nil ID, which should never happen: %s", detail)
|
||||
}
|
||||
groupIDsToRemove = append(groupIDsToRemove, *r.GroupID)
|
||||
}
|
||||
|
||||
// HandleMissingGroups will add the new groups to the org if
|
||||
// the settings specify. It will convert all group names into uuids
|
||||
// for easier assignment.
|
||||
// TODO: This code should be batched at the end of the for loop.
|
||||
// Optimizing this is being pushed because if AutoCreate is disabled,
|
||||
// this code will only add cost on the first login for each user.
|
||||
// AutoCreate is usually disabled for large deployments.
|
||||
// For small deployments, this is less of a problem.
|
||||
assignGroups, err := settings.HandleMissingGroups(ctx, tx, orgID, add)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("handle missing groups: %w", err)
|
||||
}
|
||||
|
||||
groupIDsToAdd = append(groupIDsToAdd, assignGroups...)
|
||||
}
|
||||
|
||||
// ApplyGroupDifference will take the total adds and removes, and apply
|
||||
// them.
|
||||
err = s.ApplyGroupDifference(ctx, tx, user, groupIDsToAdd, groupIDsToRemove)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("apply group difference: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ApplyGroupDifference will add and remove the user from the specified groups.
|
||||
func (s AGPLIDPSync) ApplyGroupDifference(ctx context.Context, tx database.Store, user database.User, add []uuid.UUID, removeIDs []uuid.UUID) error {
|
||||
if len(removeIDs) > 0 {
|
||||
removedGroupIDs, err := tx.RemoveUserFromGroups(ctx, database.RemoveUserFromGroupsParams{
|
||||
UserID: user.ID,
|
||||
GroupIds: removeIDs,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("remove user from %d groups: %w", len(removeIDs), err)
|
||||
}
|
||||
if len(removedGroupIDs) != len(removeIDs) {
|
||||
s.Logger.Debug(ctx, "user not removed from expected number of groups",
|
||||
slog.F("user_id", user.ID),
|
||||
slog.F("groups_removed_count", len(removedGroupIDs)),
|
||||
slog.F("expected_count", len(removeIDs)),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
if len(add) > 0 {
|
||||
add = slice.Unique(add)
|
||||
// Defensive programming to only insert uniques.
|
||||
assignedGroupIDs, err := tx.InsertUserGroupsByID(ctx, database.InsertUserGroupsByIDParams{
|
||||
UserID: user.ID,
|
||||
GroupIds: add,
|
||||
})
|
||||
if err != nil {
|
||||
return xerrors.Errorf("insert user into %d groups: %w", len(add), err)
|
||||
}
|
||||
if len(assignedGroupIDs) != len(add) {
|
||||
s.Logger.Debug(ctx, "user not assigned to expected number of groups",
|
||||
slog.F("user_id", user.ID),
|
||||
slog.F("groups_assigned_count", len(assignedGroupIDs)),
|
||||
slog.F("expected_count", len(add)),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
type GroupSyncSettings struct {
|
||||
// Field selects the claim field to be used as the created user's
|
||||
// groups. If the group field is the empty string, then no group updates
|
||||
// will ever come from the OIDC provider.
|
||||
Field string `json:"field"`
|
||||
// Mapping maps from an OIDC group --> Coder group ID
|
||||
Mapping map[string][]uuid.UUID `json:"mapping"`
|
||||
// RegexFilter is a regular expression that filters the groups returned by
|
||||
// the OIDC provider. Any group not matched by this regex will be ignored.
|
||||
// If the group filter is nil, then no group filtering will occur.
|
||||
RegexFilter *regexp.Regexp `json:"regex_filter"`
|
||||
// AutoCreateMissing controls whether groups returned by the OIDC provider
|
||||
// are automatically created in Coder if they are missing.
|
||||
AutoCreateMissing bool `json:"auto_create_missing_groups"`
|
||||
// LegacyNameMapping is deprecated. It remaps an IDP group name to
|
||||
// a Coder group name. Since configuration is now done at runtime,
|
||||
// group IDs are used to account for group renames.
|
||||
// For legacy configurations, this config option has to remain.
|
||||
// Deprecated: Use Mapping instead.
|
||||
LegacyNameMapping map[string]string `json:"legacy_group_name_mapping,omitempty"`
|
||||
}
|
||||
|
||||
func (s *GroupSyncSettings) Set(v string) error {
|
||||
return json.Unmarshal([]byte(v), s)
|
||||
}
|
||||
|
||||
func (s *GroupSyncSettings) String() string {
|
||||
return runtimeconfig.JSONString(s)
|
||||
}
|
||||
|
||||
type ExpectedGroup struct {
|
||||
OrganizationID uuid.UUID
|
||||
GroupID *uuid.UUID
|
||||
GroupName *string
|
||||
}
|
||||
|
||||
// Equal compares two ExpectedGroups. The org id must be the same.
|
||||
// If the group ID is set, it will be compared and take priority, ignoring the
|
||||
// name value. So 2 groups with the same ID but different names will be
|
||||
// considered equal.
|
||||
func (a ExpectedGroup) Equal(b ExpectedGroup) bool {
|
||||
// Must match
|
||||
if a.OrganizationID != b.OrganizationID {
|
||||
return false
|
||||
}
|
||||
// Only the name or the name needs to be checked, priority is given to the ID.
|
||||
if a.GroupID != nil && b.GroupID != nil {
|
||||
return *a.GroupID == *b.GroupID
|
||||
}
|
||||
if a.GroupName != nil && b.GroupName != nil {
|
||||
return *a.GroupName == *b.GroupName
|
||||
}
|
||||
|
||||
// If everything is nil, it is equal. Although a bit pointless
|
||||
if a.GroupID == nil && b.GroupID == nil &&
|
||||
a.GroupName == nil && b.GroupName == nil {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// ParseClaims will take the merged claims from the IDP and return the groups
|
||||
// the user is expected to be a member of. The expected group can either be a
|
||||
// name or an ID.
|
||||
// It is unfortunate we cannot use exclusively names or exclusively IDs.
|
||||
// When configuring though, if a group is mapped from "A" -> "UUID 1234", and
|
||||
// the group "UUID 1234" is renamed, we want to maintain the mapping.
|
||||
// We have to keep names because group sync supports syncing groups by name if
|
||||
// the external IDP group name matches the Coder one.
|
||||
func (s GroupSyncSettings) ParseClaims(orgID uuid.UUID, mergedClaims jwt.MapClaims) ([]ExpectedGroup, error) {
|
||||
groupsRaw, ok := mergedClaims[s.Field]
|
||||
if !ok {
|
||||
return []ExpectedGroup{}, nil
|
||||
}
|
||||
|
||||
parsedGroups, err := ParseStringSliceClaim(groupsRaw)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("parse groups field, unexpected type %T: %w", groupsRaw, err)
|
||||
}
|
||||
|
||||
groups := make([]ExpectedGroup, 0)
|
||||
for _, group := range parsedGroups {
|
||||
group := group
|
||||
|
||||
// Legacy group mappings happen before the regex filter.
|
||||
mappedGroupName, ok := s.LegacyNameMapping[group]
|
||||
if ok {
|
||||
group = mappedGroupName
|
||||
}
|
||||
|
||||
// Only allow through groups that pass the regex
|
||||
if s.RegexFilter != nil {
|
||||
if !s.RegexFilter.MatchString(group) {
|
||||
continue
|
||||
}
|
||||
}
|
||||
|
||||
mappedGroupIDs, ok := s.Mapping[group]
|
||||
if ok {
|
||||
for _, gid := range mappedGroupIDs {
|
||||
gid := gid
|
||||
groups = append(groups, ExpectedGroup{OrganizationID: orgID, GroupID: &gid})
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
groups = append(groups, ExpectedGroup{OrganizationID: orgID, GroupName: &group})
|
||||
}
|
||||
|
||||
return groups, nil
|
||||
}
|
||||
|
||||
// HandleMissingGroups ensures all ExpectedGroups convert to uuids.
|
||||
// Groups can be referenced by name via legacy params or IDP group names.
|
||||
// These group names are converted to IDs for easier assignment.
|
||||
// Missing groups are created if AutoCreate is enabled.
|
||||
// TODO: Batching this would be better, as this is 1 or 2 db calls per organization.
|
||||
func (s GroupSyncSettings) HandleMissingGroups(ctx context.Context, tx database.Store, orgID uuid.UUID, add []ExpectedGroup) ([]uuid.UUID, error) {
|
||||
// All expected that are missing IDs means the group does not exist
|
||||
// in the database, or it is a legacy mapping, and we need to do a lookup.
|
||||
var missingGroups []string
|
||||
addIDs := make([]uuid.UUID, 0)
|
||||
|
||||
for _, expected := range add {
|
||||
if expected.GroupID == nil && expected.GroupName != nil {
|
||||
missingGroups = append(missingGroups, *expected.GroupName)
|
||||
} else if expected.GroupID != nil {
|
||||
// Keep the IDs to sync the groups.
|
||||
addIDs = append(addIDs, *expected.GroupID)
|
||||
}
|
||||
}
|
||||
|
||||
if s.AutoCreateMissing && len(missingGroups) > 0 {
|
||||
// Insert any missing groups. If the groups already exist, this is a noop.
|
||||
_, err := tx.InsertMissingGroups(ctx, database.InsertMissingGroupsParams{
|
||||
OrganizationID: orgID,
|
||||
Source: database.GroupSourceOidc,
|
||||
GroupNames: missingGroups,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("insert missing groups: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Fetch any missing groups by name. If they exist, their IDs will be
|
||||
// matched and returned.
|
||||
if len(missingGroups) > 0 {
|
||||
// Do name lookups for all groups that are missing IDs.
|
||||
newGroups, err := tx.GetGroups(ctx, database.GetGroupsParams{
|
||||
OrganizationID: orgID,
|
||||
HasMemberID: uuid.UUID{},
|
||||
GroupNames: missingGroups,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("get groups by names: %w", err)
|
||||
}
|
||||
for _, g := range newGroups {
|
||||
addIDs = append(addIDs, g.Group.ID)
|
||||
}
|
||||
}
|
||||
|
||||
return addIDs, nil
|
||||
}
|
||||
|
||||
func ConvertAllowList(allowList []string) map[string]struct{} {
|
||||
allowMap := make(map[string]struct{}, len(allowList))
|
||||
for _, group := range allowList {
|
||||
allowMap[group] = struct{}{}
|
||||
}
|
||||
return allowMap
|
||||
}
|
||||
@@ -0,0 +1,814 @@
|
||||
package idpsync_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"regexp"
|
||||
"testing"
|
||||
|
||||
"github.com/golang-jwt/jwt/v4"
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/exp/slices"
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/coderd/coderdtest"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/coderd/database/dbauthz"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/idpsync"
|
||||
"github.com/coder/coder/v2/coderd/runtimeconfig"
|
||||
"github.com/coder/coder/v2/coderd/util/ptr"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
func TestParseGroupClaims(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("EmptyConfig", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}),
|
||||
runtimeconfig.NewManager(),
|
||||
idpsync.DeploymentSyncSettings{})
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
|
||||
params, err := s.ParseGroupClaims(ctx, jwt.MapClaims{})
|
||||
require.Nil(t, err)
|
||||
|
||||
require.False(t, params.SyncEnabled)
|
||||
})
|
||||
|
||||
// AllowList has no effect in AGPL
|
||||
t.Run("AllowList", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}),
|
||||
runtimeconfig.NewManager(),
|
||||
idpsync.DeploymentSyncSettings{
|
||||
GroupField: "groups",
|
||||
GroupAllowList: map[string]struct{}{
|
||||
"foo": {},
|
||||
},
|
||||
})
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
|
||||
params, err := s.ParseGroupClaims(ctx, jwt.MapClaims{})
|
||||
require.Nil(t, err)
|
||||
require.False(t, params.SyncEnabled)
|
||||
})
|
||||
}
|
||||
|
||||
func TestGroupSyncTable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Last checked, takes 30s with postgres on a fast machine.
|
||||
if dbtestutil.WillUsePostgres() {
|
||||
t.Skip("Skipping test because it populates a lot of db entries, which is slow on postgres.")
|
||||
}
|
||||
|
||||
userClaims := jwt.MapClaims{
|
||||
"groups": []string{
|
||||
"foo", "bar", "baz",
|
||||
"create-bar", "create-baz",
|
||||
"legacy-bar",
|
||||
},
|
||||
}
|
||||
|
||||
ids := coderdtest.NewDeterministicUUIDGenerator()
|
||||
testCases := []orgSetupDefinition{
|
||||
{
|
||||
Name: "SwitchGroups",
|
||||
Settings: &idpsync.GroupSyncSettings{
|
||||
Field: "groups",
|
||||
Mapping: map[string][]uuid.UUID{
|
||||
"foo": {ids.ID("sg-foo"), ids.ID("sg-foo-2")},
|
||||
"bar": {ids.ID("sg-bar")},
|
||||
"baz": {ids.ID("sg-baz")},
|
||||
},
|
||||
},
|
||||
Groups: map[uuid.UUID]bool{
|
||||
uuid.New(): true,
|
||||
uuid.New(): true,
|
||||
// Extra groups
|
||||
ids.ID("sg-foo"): false,
|
||||
ids.ID("sg-foo-2"): false,
|
||||
ids.ID("sg-bar"): false,
|
||||
ids.ID("sg-baz"): false,
|
||||
},
|
||||
ExpectedGroups: []uuid.UUID{
|
||||
ids.ID("sg-foo"),
|
||||
ids.ID("sg-foo-2"),
|
||||
ids.ID("sg-bar"),
|
||||
ids.ID("sg-baz"),
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "StayInGroup",
|
||||
Settings: &idpsync.GroupSyncSettings{
|
||||
Field: "groups",
|
||||
// Only match foo, so bar does not map
|
||||
RegexFilter: regexp.MustCompile("^foo$"),
|
||||
Mapping: map[string][]uuid.UUID{
|
||||
"foo": {ids.ID("gg-foo"), uuid.New()},
|
||||
"bar": {ids.ID("gg-bar")},
|
||||
"baz": {ids.ID("gg-baz")},
|
||||
},
|
||||
},
|
||||
Groups: map[uuid.UUID]bool{
|
||||
ids.ID("gg-foo"): true,
|
||||
ids.ID("gg-bar"): false,
|
||||
},
|
||||
ExpectedGroups: []uuid.UUID{
|
||||
ids.ID("gg-foo"),
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "UserJoinsGroups",
|
||||
Settings: &idpsync.GroupSyncSettings{
|
||||
Field: "groups",
|
||||
Mapping: map[string][]uuid.UUID{
|
||||
"foo": {ids.ID("ng-foo"), uuid.New()},
|
||||
"bar": {ids.ID("ng-bar"), ids.ID("ng-bar-2")},
|
||||
"baz": {ids.ID("ng-baz")},
|
||||
},
|
||||
},
|
||||
Groups: map[uuid.UUID]bool{
|
||||
ids.ID("ng-foo"): false,
|
||||
ids.ID("ng-bar"): false,
|
||||
ids.ID("ng-bar-2"): false,
|
||||
ids.ID("ng-baz"): false,
|
||||
},
|
||||
ExpectedGroups: []uuid.UUID{
|
||||
ids.ID("ng-foo"),
|
||||
ids.ID("ng-bar"),
|
||||
ids.ID("ng-bar-2"),
|
||||
ids.ID("ng-baz"),
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "CreateGroups",
|
||||
Settings: &idpsync.GroupSyncSettings{
|
||||
Field: "groups",
|
||||
RegexFilter: regexp.MustCompile("^create"),
|
||||
AutoCreateMissing: true,
|
||||
},
|
||||
Groups: map[uuid.UUID]bool{},
|
||||
ExpectedGroupNames: []string{
|
||||
"create-bar",
|
||||
"create-baz",
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "GroupNamesNoMapping",
|
||||
Settings: &idpsync.GroupSyncSettings{
|
||||
Field: "groups",
|
||||
RegexFilter: regexp.MustCompile(".*"),
|
||||
AutoCreateMissing: false,
|
||||
},
|
||||
GroupNames: map[string]bool{
|
||||
"foo": false,
|
||||
"bar": false,
|
||||
"goob": true,
|
||||
},
|
||||
ExpectedGroupNames: []string{
|
||||
"foo",
|
||||
"bar",
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "NoUser",
|
||||
Settings: &idpsync.GroupSyncSettings{
|
||||
Field: "groups",
|
||||
Mapping: map[string][]uuid.UUID{
|
||||
// Extra ID that does not map to a group
|
||||
"foo": {ids.ID("ow-foo"), uuid.New()},
|
||||
},
|
||||
RegexFilter: nil,
|
||||
AutoCreateMissing: false,
|
||||
},
|
||||
NotMember: true,
|
||||
Groups: map[uuid.UUID]bool{
|
||||
ids.ID("ow-foo"): false,
|
||||
ids.ID("ow-bar"): false,
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "NoSettingsNoUser",
|
||||
Settings: nil,
|
||||
Groups: map[uuid.UUID]bool{},
|
||||
},
|
||||
{
|
||||
Name: "LegacyMapping",
|
||||
Settings: &idpsync.GroupSyncSettings{
|
||||
Field: "groups",
|
||||
RegexFilter: regexp.MustCompile("^legacy"),
|
||||
LegacyNameMapping: map[string]string{
|
||||
"create-bar": "legacy-bar",
|
||||
"foo": "legacy-foo",
|
||||
"bop": "legacy-bop",
|
||||
},
|
||||
AutoCreateMissing: true,
|
||||
},
|
||||
Groups: map[uuid.UUID]bool{
|
||||
ids.ID("lg-foo"): true,
|
||||
},
|
||||
GroupNames: map[string]bool{
|
||||
"legacy-foo": false,
|
||||
"extra": true,
|
||||
"legacy-bop": true,
|
||||
},
|
||||
ExpectedGroupNames: []string{
|
||||
"legacy-bar",
|
||||
"legacy-foo",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
tc := tc
|
||||
t.Run(tc.Name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
manager := runtimeconfig.NewManager()
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}),
|
||||
manager,
|
||||
idpsync.DeploymentSyncSettings{
|
||||
GroupField: "groups",
|
||||
},
|
||||
)
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitSuperLong)
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
orgID := uuid.New()
|
||||
SetupOrganization(t, s, db, user, orgID, tc)
|
||||
|
||||
// Do the group sync!
|
||||
err := s.SyncGroups(ctx, db, user, idpsync.GroupParams{
|
||||
SyncEnabled: true,
|
||||
MergedClaims: userClaims,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
tc.Assert(t, orgID, db, user)
|
||||
})
|
||||
}
|
||||
|
||||
// AllTogether runs the entire tabled test as a singular user and
|
||||
// deployment. This tests all organizations being synced together.
|
||||
// The reason we do them individually, is that it is much easier to
|
||||
// debug a single test case.
|
||||
t.Run("AllTogether", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
manager := runtimeconfig.NewManager()
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}),
|
||||
manager,
|
||||
// Also sync the default org!
|
||||
idpsync.DeploymentSyncSettings{
|
||||
GroupField: "groups",
|
||||
Legacy: idpsync.DefaultOrgLegacySettings{
|
||||
GroupField: "groups",
|
||||
GroupMapping: map[string]string{
|
||||
"foo": "legacy-foo",
|
||||
"baz": "legacy-baz",
|
||||
},
|
||||
GroupFilter: regexp.MustCompile("^legacy"),
|
||||
CreateMissingGroups: true,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitSuperLong)
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
|
||||
var asserts []func(t *testing.T)
|
||||
// The default org is also going to do something
|
||||
def := orgSetupDefinition{
|
||||
Name: "DefaultOrg",
|
||||
GroupNames: map[string]bool{
|
||||
"legacy-foo": false,
|
||||
"legacy-baz": true,
|
||||
"random": true,
|
||||
},
|
||||
// No settings, because they come from the deployment values
|
||||
Settings: nil,
|
||||
ExpectedGroups: nil,
|
||||
ExpectedGroupNames: []string{"legacy-foo", "legacy-baz", "legacy-bar"},
|
||||
}
|
||||
|
||||
//nolint:gocritic // testing
|
||||
defOrg, err := db.GetDefaultOrganization(dbauthz.AsSystemRestricted(ctx))
|
||||
require.NoError(t, err)
|
||||
SetupOrganization(t, s, db, user, defOrg.ID, def)
|
||||
asserts = append(asserts, func(t *testing.T) {
|
||||
t.Run(def.Name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
def.Assert(t, defOrg.ID, db, user)
|
||||
})
|
||||
})
|
||||
|
||||
for _, tc := range testCases {
|
||||
tc := tc
|
||||
|
||||
orgID := uuid.New()
|
||||
SetupOrganization(t, s, db, user, orgID, tc)
|
||||
asserts = append(asserts, func(t *testing.T) {
|
||||
t.Run(tc.Name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
tc.Assert(t, orgID, db, user)
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
asserts = append(asserts, func(t *testing.T) {
|
||||
t.Helper()
|
||||
def.Assert(t, defOrg.ID, db, user)
|
||||
})
|
||||
|
||||
// Do the group sync!
|
||||
err = s.SyncGroups(ctx, db, user, idpsync.GroupParams{
|
||||
SyncEnabled: true,
|
||||
MergedClaims: userClaims,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
for _, assert := range asserts {
|
||||
assert(t)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSyncDisabled(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
if dbtestutil.WillUsePostgres() {
|
||||
t.Skip("Skipping test because it populates a lot of db entries, which is slow on postgres.")
|
||||
}
|
||||
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
manager := runtimeconfig.NewManager()
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}),
|
||||
manager,
|
||||
idpsync.DeploymentSyncSettings{},
|
||||
)
|
||||
|
||||
ids := coderdtest.NewDeterministicUUIDGenerator()
|
||||
ctx := testutil.Context(t, testutil.WaitSuperLong)
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
orgID := uuid.New()
|
||||
|
||||
def := orgSetupDefinition{
|
||||
Name: "SyncDisabled",
|
||||
Groups: map[uuid.UUID]bool{
|
||||
ids.ID("foo"): true,
|
||||
ids.ID("bar"): true,
|
||||
ids.ID("baz"): false,
|
||||
ids.ID("bop"): false,
|
||||
},
|
||||
Settings: &idpsync.GroupSyncSettings{
|
||||
Field: "groups",
|
||||
Mapping: map[string][]uuid.UUID{
|
||||
"foo": {ids.ID("foo")},
|
||||
"baz": {ids.ID("baz")},
|
||||
},
|
||||
},
|
||||
ExpectedGroups: []uuid.UUID{
|
||||
ids.ID("foo"),
|
||||
ids.ID("bar"),
|
||||
},
|
||||
}
|
||||
|
||||
SetupOrganization(t, s, db, user, orgID, def)
|
||||
|
||||
// Do the group sync!
|
||||
err := s.SyncGroups(ctx, db, user, idpsync.GroupParams{
|
||||
SyncEnabled: false,
|
||||
MergedClaims: jwt.MapClaims{
|
||||
"groups": []string{"baz", "bop"},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
def.Assert(t, orgID, db, user)
|
||||
}
|
||||
|
||||
// TestApplyGroupDifference is mainly testing the database functions
|
||||
func TestApplyGroupDifference(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ids := coderdtest.NewDeterministicUUIDGenerator()
|
||||
testCase := []struct {
|
||||
Name string
|
||||
Before map[uuid.UUID]bool
|
||||
Add []uuid.UUID
|
||||
Remove []uuid.UUID
|
||||
Expect []uuid.UUID
|
||||
}{
|
||||
{
|
||||
Name: "Empty",
|
||||
},
|
||||
{
|
||||
Name: "AddFromNone",
|
||||
Before: map[uuid.UUID]bool{
|
||||
ids.ID("g1"): false,
|
||||
},
|
||||
Add: []uuid.UUID{
|
||||
ids.ID("g1"),
|
||||
},
|
||||
Expect: []uuid.UUID{
|
||||
ids.ID("g1"),
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "AddSome",
|
||||
Before: map[uuid.UUID]bool{
|
||||
ids.ID("g1"): true,
|
||||
ids.ID("g2"): false,
|
||||
ids.ID("g3"): false,
|
||||
uuid.New(): false,
|
||||
},
|
||||
Add: []uuid.UUID{
|
||||
ids.ID("g2"),
|
||||
ids.ID("g3"),
|
||||
},
|
||||
Expect: []uuid.UUID{
|
||||
ids.ID("g1"),
|
||||
ids.ID("g2"),
|
||||
ids.ID("g3"),
|
||||
},
|
||||
},
|
||||
{
|
||||
Name: "RemoveAll",
|
||||
Before: map[uuid.UUID]bool{
|
||||
uuid.New(): false,
|
||||
ids.ID("g2"): true,
|
||||
ids.ID("g3"): true,
|
||||
},
|
||||
Remove: []uuid.UUID{
|
||||
ids.ID("g2"),
|
||||
ids.ID("g3"),
|
||||
},
|
||||
Expect: []uuid.UUID{},
|
||||
},
|
||||
{
|
||||
Name: "Mixed",
|
||||
Before: map[uuid.UUID]bool{
|
||||
// adds
|
||||
ids.ID("a1"): true,
|
||||
ids.ID("a2"): true,
|
||||
ids.ID("a3"): false,
|
||||
ids.ID("a4"): false,
|
||||
// removes
|
||||
ids.ID("r1"): true,
|
||||
ids.ID("r2"): true,
|
||||
ids.ID("r3"): false,
|
||||
ids.ID("r4"): false,
|
||||
// stable
|
||||
ids.ID("s1"): true,
|
||||
ids.ID("s2"): true,
|
||||
// noise
|
||||
uuid.New(): false,
|
||||
uuid.New(): false,
|
||||
},
|
||||
Add: []uuid.UUID{
|
||||
ids.ID("a1"), ids.ID("a2"),
|
||||
ids.ID("a3"), ids.ID("a4"),
|
||||
// Double up to try and confuse
|
||||
ids.ID("a1"),
|
||||
ids.ID("a4"),
|
||||
},
|
||||
Remove: []uuid.UUID{
|
||||
ids.ID("r1"), ids.ID("r2"),
|
||||
ids.ID("r3"), ids.ID("r4"),
|
||||
// Double up to try and confuse
|
||||
ids.ID("r1"),
|
||||
ids.ID("r4"),
|
||||
},
|
||||
Expect: []uuid.UUID{
|
||||
ids.ID("a1"), ids.ID("a2"), ids.ID("a3"), ids.ID("a4"),
|
||||
ids.ID("s1"), ids.ID("s2"),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCase {
|
||||
tc := tc
|
||||
t.Run(tc.Name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
mgr := runtimeconfig.NewManager()
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
//nolint:gocritic // testing
|
||||
ctx = dbauthz.AsSystemRestricted(ctx)
|
||||
|
||||
org := dbgen.Organization(t, db, database.Organization{})
|
||||
_, err := db.InsertAllUsersGroup(ctx, org.ID)
|
||||
require.NoError(t, err)
|
||||
|
||||
user := dbgen.User(t, db, database.User{})
|
||||
_ = dbgen.OrganizationMember(t, db, database.OrganizationMember{
|
||||
UserID: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
|
||||
for gid, in := range tc.Before {
|
||||
group := dbgen.Group(t, db, database.Group{
|
||||
ID: gid,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
if in {
|
||||
_ = dbgen.GroupMember(t, db, database.GroupMemberTable{
|
||||
UserID: user.ID,
|
||||
GroupID: group.ID,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}), mgr, idpsync.FromDeploymentValues(coderdtest.DeploymentValues(t)))
|
||||
err = s.ApplyGroupDifference(context.Background(), db, user, tc.Add, tc.Remove)
|
||||
require.NoError(t, err)
|
||||
|
||||
userGroups, err := db.GetGroups(ctx, database.GetGroupsParams{
|
||||
HasMemberID: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// assert
|
||||
found := db2sdk.List(userGroups, func(g database.GetGroupsRow) uuid.UUID {
|
||||
return g.Group.ID
|
||||
})
|
||||
|
||||
// Add everyone group
|
||||
require.ElementsMatch(t, append(tc.Expect, org.ID), found)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExpectedGroupEqual(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ids := coderdtest.NewDeterministicUUIDGenerator()
|
||||
testCases := []struct {
|
||||
Name string
|
||||
A idpsync.ExpectedGroup
|
||||
B idpsync.ExpectedGroup
|
||||
Equal bool
|
||||
}{
|
||||
{
|
||||
Name: "Empty",
|
||||
A: idpsync.ExpectedGroup{},
|
||||
B: idpsync.ExpectedGroup{},
|
||||
Equal: true,
|
||||
},
|
||||
{
|
||||
Name: "DifferentOrgs",
|
||||
A: idpsync.ExpectedGroup{
|
||||
OrganizationID: uuid.New(),
|
||||
GroupID: ptr.Ref(ids.ID("g1")),
|
||||
GroupName: nil,
|
||||
},
|
||||
B: idpsync.ExpectedGroup{
|
||||
OrganizationID: uuid.New(),
|
||||
GroupID: ptr.Ref(ids.ID("g1")),
|
||||
GroupName: nil,
|
||||
},
|
||||
Equal: false,
|
||||
},
|
||||
{
|
||||
Name: "SameID",
|
||||
A: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: ptr.Ref(ids.ID("g1")),
|
||||
GroupName: nil,
|
||||
},
|
||||
B: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: ptr.Ref(ids.ID("g1")),
|
||||
GroupName: nil,
|
||||
},
|
||||
Equal: true,
|
||||
},
|
||||
{
|
||||
Name: "DifferentIDs",
|
||||
A: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: ptr.Ref(uuid.New()),
|
||||
GroupName: nil,
|
||||
},
|
||||
B: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: ptr.Ref(uuid.New()),
|
||||
GroupName: nil,
|
||||
},
|
||||
Equal: false,
|
||||
},
|
||||
{
|
||||
Name: "SameName",
|
||||
A: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: nil,
|
||||
GroupName: ptr.Ref("foo"),
|
||||
},
|
||||
B: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: nil,
|
||||
GroupName: ptr.Ref("foo"),
|
||||
},
|
||||
Equal: true,
|
||||
},
|
||||
{
|
||||
Name: "DifferentName",
|
||||
A: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: nil,
|
||||
GroupName: ptr.Ref("foo"),
|
||||
},
|
||||
B: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: nil,
|
||||
GroupName: ptr.Ref("bar"),
|
||||
},
|
||||
Equal: false,
|
||||
},
|
||||
// Edge cases
|
||||
{
|
||||
// A bit strange, but valid as ID takes priority.
|
||||
// We assume 2 groups with the same ID are equal, even if
|
||||
// their names are different. Names are mutable, IDs are not,
|
||||
// so there is 0% chance they are different groups.
|
||||
Name: "DifferentIDSameName",
|
||||
A: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: ptr.Ref(ids.ID("g1")),
|
||||
GroupName: ptr.Ref("foo"),
|
||||
},
|
||||
B: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: ptr.Ref(ids.ID("g1")),
|
||||
GroupName: ptr.Ref("bar"),
|
||||
},
|
||||
Equal: true,
|
||||
},
|
||||
{
|
||||
Name: "MixedNils",
|
||||
A: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: ptr.Ref(ids.ID("g1")),
|
||||
GroupName: nil,
|
||||
},
|
||||
B: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: nil,
|
||||
GroupName: ptr.Ref("bar"),
|
||||
},
|
||||
Equal: false,
|
||||
},
|
||||
{
|
||||
Name: "NoComparable",
|
||||
A: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: ptr.Ref(ids.ID("g1")),
|
||||
GroupName: nil,
|
||||
},
|
||||
B: idpsync.ExpectedGroup{
|
||||
OrganizationID: ids.ID("org"),
|
||||
GroupID: nil,
|
||||
GroupName: nil,
|
||||
},
|
||||
Equal: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
tc := tc
|
||||
t.Run(tc.Name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.Equal(t, tc.Equal, tc.A.Equal(tc.B))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func SetupOrganization(t *testing.T, s *idpsync.AGPLIDPSync, db database.Store, user database.User, orgID uuid.UUID, def orgSetupDefinition) {
|
||||
t.Helper()
|
||||
|
||||
// Account that the org might be the default organization
|
||||
org, err := db.GetOrganizationByID(context.Background(), orgID)
|
||||
if xerrors.Is(err, sql.ErrNoRows) {
|
||||
org = dbgen.Organization(t, db, database.Organization{
|
||||
ID: orgID,
|
||||
})
|
||||
}
|
||||
|
||||
_, err = db.InsertAllUsersGroup(context.Background(), org.ID)
|
||||
if !database.IsUniqueViolation(err) {
|
||||
require.NoError(t, err, "Everyone group for an org")
|
||||
}
|
||||
|
||||
manager := runtimeconfig.NewManager()
|
||||
orgResolver := manager.OrganizationResolver(db, org.ID)
|
||||
err = s.Group.SetRuntimeValue(context.Background(), orgResolver, def.Settings)
|
||||
require.NoError(t, err)
|
||||
|
||||
if !def.NotMember {
|
||||
dbgen.OrganizationMember(t, db, database.OrganizationMember{
|
||||
UserID: user.ID,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
}
|
||||
for groupID, in := range def.Groups {
|
||||
dbgen.Group(t, db, database.Group{
|
||||
ID: groupID,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
if in {
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{
|
||||
UserID: user.ID,
|
||||
GroupID: groupID,
|
||||
})
|
||||
}
|
||||
}
|
||||
for groupName, in := range def.GroupNames {
|
||||
group := dbgen.Group(t, db, database.Group{
|
||||
Name: groupName,
|
||||
OrganizationID: org.ID,
|
||||
})
|
||||
if in {
|
||||
dbgen.GroupMember(t, db, database.GroupMemberTable{
|
||||
UserID: user.ID,
|
||||
GroupID: group.ID,
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type orgSetupDefinition struct {
|
||||
Name string
|
||||
// True if the user is a member of the group
|
||||
Groups map[uuid.UUID]bool
|
||||
GroupNames map[string]bool
|
||||
NotMember bool
|
||||
|
||||
Settings *idpsync.GroupSyncSettings
|
||||
ExpectedGroups []uuid.UUID
|
||||
ExpectedGroupNames []string
|
||||
}
|
||||
|
||||
func (o orgSetupDefinition) Assert(t *testing.T, orgID uuid.UUID, db database.Store, user database.User) {
|
||||
t.Helper()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
members, err := db.OrganizationMembers(ctx, database.OrganizationMembersParams{
|
||||
OrganizationID: orgID,
|
||||
UserID: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
if o.NotMember {
|
||||
require.Len(t, members, 0, "should not be a member")
|
||||
} else {
|
||||
require.Len(t, members, 1, "should be a member")
|
||||
}
|
||||
|
||||
userGroups, err := db.GetGroups(ctx, database.GetGroupsParams{
|
||||
OrganizationID: orgID,
|
||||
HasMemberID: user.ID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
if o.ExpectedGroups == nil {
|
||||
o.ExpectedGroups = make([]uuid.UUID, 0)
|
||||
}
|
||||
if len(o.ExpectedGroupNames) > 0 && len(o.ExpectedGroups) > 0 {
|
||||
t.Fatal("ExpectedGroups and ExpectedGroupNames are mutually exclusive")
|
||||
}
|
||||
|
||||
// Everyone groups mess up our asserts
|
||||
userGroups = slices.DeleteFunc(userGroups, func(row database.GetGroupsRow) bool {
|
||||
return row.Group.ID == row.Group.OrganizationID
|
||||
})
|
||||
|
||||
if len(o.ExpectedGroupNames) > 0 {
|
||||
found := db2sdk.List(userGroups, func(g database.GetGroupsRow) string {
|
||||
return g.Group.Name
|
||||
})
|
||||
require.ElementsMatch(t, o.ExpectedGroupNames, found, "user groups by name")
|
||||
require.Len(t, o.ExpectedGroups, 0, "ExpectedGroups should be empty")
|
||||
} else {
|
||||
// Check by ID, recommended
|
||||
found := db2sdk.List(userGroups, func(g database.GetGroupsRow) uuid.UUID {
|
||||
return g.Group.ID
|
||||
})
|
||||
require.ElementsMatch(t, o.ExpectedGroups, found, "user groups")
|
||||
require.Len(t, o.ExpectedGroupNames, 0, "ExpectedGroupNames should be empty")
|
||||
}
|
||||
}
|
||||
+69
-15
@@ -3,6 +3,7 @@ package idpsync
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/golang-jwt/jwt/v4"
|
||||
@@ -12,6 +13,7 @@ import (
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/httpapi"
|
||||
"github.com/coder/coder/v2/coderd/runtimeconfig"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/coder/v2/site"
|
||||
)
|
||||
@@ -25,21 +27,34 @@ type IDPSync interface {
|
||||
OrganizationSyncEnabled() bool
|
||||
// ParseOrganizationClaims takes claims from an OIDC provider, and returns the
|
||||
// organization sync params for assigning users into organizations.
|
||||
ParseOrganizationClaims(ctx context.Context, _ jwt.MapClaims) (OrganizationParams, *HTTPError)
|
||||
ParseOrganizationClaims(ctx context.Context, mergedClaims jwt.MapClaims) (OrganizationParams, *HTTPError)
|
||||
// SyncOrganizations assigns and removed users from organizations based on the
|
||||
// provided params.
|
||||
SyncOrganizations(ctx context.Context, tx database.Store, user database.User, params OrganizationParams) error
|
||||
|
||||
GroupSyncEnabled() bool
|
||||
// ParseGroupClaims takes claims from an OIDC provider, and returns the params
|
||||
// for group syncing. Most of the logic happens in SyncGroups.
|
||||
ParseGroupClaims(ctx context.Context, mergedClaims jwt.MapClaims) (GroupParams, *HTTPError)
|
||||
// SyncGroups assigns and removes users from groups based on the provided params.
|
||||
SyncGroups(ctx context.Context, db database.Store, user database.User, params GroupParams) error
|
||||
// GroupSyncSettings is exposed for the API to implement CRUD operations
|
||||
// on the settings used by IDPSync. This entry is thread safe and can be
|
||||
// accessed concurrently. The settings are stored in the database.
|
||||
GroupSyncSettings() runtimeconfig.RuntimeEntry[*GroupSyncSettings]
|
||||
}
|
||||
|
||||
// AGPLIDPSync is the configuration for syncing user information from an external
|
||||
// IDP. All related code to syncing user information should be in this package.
|
||||
type AGPLIDPSync struct {
|
||||
Logger slog.Logger
|
||||
Logger slog.Logger
|
||||
Manager *runtimeconfig.Manager
|
||||
|
||||
SyncSettings
|
||||
}
|
||||
|
||||
type SyncSettings struct {
|
||||
// DeploymentSyncSettings are static and are sourced from the deployment config.
|
||||
type DeploymentSyncSettings struct {
|
||||
// OrganizationField selects the claim field to be used as the created user's
|
||||
// organizations. If the field is the empty string, then no organization updates
|
||||
// will ever come from the OIDC provider.
|
||||
@@ -50,23 +65,62 @@ type SyncSettings struct {
|
||||
// placed into the default organization. This is mostly a hack to support
|
||||
// legacy deployments.
|
||||
OrganizationAssignDefault bool
|
||||
|
||||
// GroupField at the deployment level is used for deployment level group claim
|
||||
// settings.
|
||||
GroupField string
|
||||
// GroupAllowList (if set) will restrict authentication to only users who
|
||||
// have at least one group in this list.
|
||||
// A map representation is used for easier lookup.
|
||||
GroupAllowList map[string]struct{}
|
||||
// Legacy deployment settings that only apply to the default org.
|
||||
Legacy DefaultOrgLegacySettings
|
||||
}
|
||||
|
||||
type OrganizationParams struct {
|
||||
// SyncEnabled if false will skip syncing the user's organizations.
|
||||
SyncEnabled bool
|
||||
// IncludeDefault is primarily for single org deployments. It will ensure
|
||||
// a user is always inserted into the default org.
|
||||
IncludeDefault bool
|
||||
// Organizations is the list of organizations the user should be a member of
|
||||
// assuming syncing is turned on.
|
||||
Organizations []uuid.UUID
|
||||
type DefaultOrgLegacySettings struct {
|
||||
GroupField string
|
||||
GroupMapping map[string]string
|
||||
GroupFilter *regexp.Regexp
|
||||
CreateMissingGroups bool
|
||||
}
|
||||
|
||||
func NewAGPLSync(logger slog.Logger, settings SyncSettings) *AGPLIDPSync {
|
||||
func FromDeploymentValues(dv *codersdk.DeploymentValues) DeploymentSyncSettings {
|
||||
if dv == nil {
|
||||
panic("Developer error: DeploymentValues should not be nil")
|
||||
}
|
||||
return DeploymentSyncSettings{
|
||||
OrganizationField: dv.OIDC.OrganizationField.Value(),
|
||||
OrganizationMapping: dv.OIDC.OrganizationMapping.Value,
|
||||
OrganizationAssignDefault: dv.OIDC.OrganizationAssignDefault.Value(),
|
||||
|
||||
// TODO: Separate group field for allow list from default org.
|
||||
// Right now you cannot disable group sync from the default org and
|
||||
// configure an allow list.
|
||||
GroupField: dv.OIDC.GroupField.Value(),
|
||||
GroupAllowList: ConvertAllowList(dv.OIDC.GroupAllowList.Value()),
|
||||
Legacy: DefaultOrgLegacySettings{
|
||||
GroupField: dv.OIDC.GroupField.Value(),
|
||||
GroupMapping: dv.OIDC.GroupMapping.Value,
|
||||
GroupFilter: dv.OIDC.GroupRegexFilter.Value(),
|
||||
CreateMissingGroups: dv.OIDC.GroupAutoCreate.Value(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
type SyncSettings struct {
|
||||
DeploymentSyncSettings
|
||||
|
||||
Group runtimeconfig.RuntimeEntry[*GroupSyncSettings]
|
||||
}
|
||||
|
||||
func NewAGPLSync(logger slog.Logger, manager *runtimeconfig.Manager, settings DeploymentSyncSettings) *AGPLIDPSync {
|
||||
return &AGPLIDPSync{
|
||||
Logger: logger.Named("idp-sync"),
|
||||
SyncSettings: settings,
|
||||
Logger: logger.Named("idp-sync"),
|
||||
Manager: manager,
|
||||
SyncSettings: SyncSettings{
|
||||
DeploymentSyncSettings: settings,
|
||||
Group: runtimeconfig.MustNew[*GroupSyncSettings]("group-sync-settings"),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -16,6 +16,17 @@ import (
|
||||
"github.com/coder/coder/v2/coderd/util/slice"
|
||||
)
|
||||
|
||||
type OrganizationParams struct {
|
||||
// SyncEnabled if false will skip syncing the user's organizations.
|
||||
SyncEnabled bool
|
||||
// IncludeDefault is primarily for single org deployments. It will ensure
|
||||
// a user is always inserted into the default org.
|
||||
IncludeDefault bool
|
||||
// Organizations is the list of organizations the user should be a member of
|
||||
// assuming syncing is turned on.
|
||||
Organizations []uuid.UUID
|
||||
}
|
||||
|
||||
func (AGPLIDPSync) OrganizationSyncEnabled() bool {
|
||||
// AGPL does not support syncing organizations.
|
||||
return false
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
"github.com/coder/coder/v2/coderd/idpsync"
|
||||
"github.com/coder/coder/v2/coderd/runtimeconfig"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
)
|
||||
|
||||
@@ -18,11 +19,13 @@ func TestParseOrganizationClaims(t *testing.T) {
|
||||
t.Run("SingleOrgDeployment", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}), idpsync.SyncSettings{
|
||||
OrganizationField: "",
|
||||
OrganizationMapping: nil,
|
||||
OrganizationAssignDefault: true,
|
||||
})
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}),
|
||||
runtimeconfig.NewManager(),
|
||||
idpsync.DeploymentSyncSettings{
|
||||
OrganizationField: "",
|
||||
OrganizationMapping: nil,
|
||||
OrganizationAssignDefault: true,
|
||||
})
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
|
||||
@@ -38,13 +41,15 @@ func TestParseOrganizationClaims(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// AGPL has limited behavior
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}), idpsync.SyncSettings{
|
||||
OrganizationField: "orgs",
|
||||
OrganizationMapping: map[string][]uuid.UUID{
|
||||
"random": {uuid.New()},
|
||||
},
|
||||
OrganizationAssignDefault: false,
|
||||
})
|
||||
s := idpsync.NewAGPLSync(slogtest.Make(t, &slogtest.Options{}),
|
||||
runtimeconfig.NewManager(),
|
||||
idpsync.DeploymentSyncSettings{
|
||||
OrganizationField: "orgs",
|
||||
OrganizationMapping: map[string][]uuid.UUID{
|
||||
"random": {uuid.New()},
|
||||
},
|
||||
OrganizationAssignDefault: false,
|
||||
})
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitMedium)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user