From 8658608c37522b6adebfa451ce9f8a33c0f78a50 Mon Sep 17 00:00:00 2001 From: Victor Lyuboslavsky Date: Wed, 2 Apr 2025 17:10:40 -0500 Subject: [PATCH] Add SCIM Groups (#27702) For #27287 This PR adds SCIM Groups to Fleet's SCIM endpoint as a follow on to SCIM Users. The logic has been manually tested with Okta, and integration tests will be in the next PR. # Checklist for submitter - [x] Added/updated automated tests - [x] Manual QA for all new/changed functionality --- ee/server/scim/groups.go | 327 ++++++++++++ ee/server/scim/scim.go | 71 ++- ee/server/scim/users.go | 55 +- server/datastore/mysql/scim.go | 476 ++++++++++++++++- server/datastore/mysql/scim_test.go | 767 +++++++++++++++++++++++++++- server/fleet/datastore.go | 12 + server/fleet/scim.go | 14 +- server/mock/datastore_mock.go | 72 +++ 8 files changed, 1757 insertions(+), 37 deletions(-) create mode 100644 ee/server/scim/groups.go diff --git a/ee/server/scim/groups.go b/ee/server/scim/groups.go new file mode 100644 index 0000000000..7c115f39da --- /dev/null +++ b/ee/server/scim/groups.go @@ -0,0 +1,327 @@ +package scim + +import ( + "fmt" + "net/http" + "strconv" + "strings" + + "github.com/elimity-com/scim" + "github.com/elimity-com/scim/errors" + "github.com/elimity-com/scim/optional" + "github.com/fleetdm/fleet/v4/server/fleet" + kitlog "github.com/go-kit/log" + "github.com/go-kit/log/level" +) + +const ( + // Group attributes: https://datatracker.ietf.org/doc/html/rfc7643#section-4.2 + displayNameAttr = "displayName" + membersAttr = "members" +) + +type GroupHandler struct { + ds fleet.Datastore + logger kitlog.Logger +} + +// Compile-time check +var _ scim.ResourceHandler = &GroupHandler{} + +func NewGroupHandler(ds fleet.Datastore, logger kitlog.Logger) scim.ResourceHandler { + return &GroupHandler{ds: ds, logger: logger} +} + +// Create creates a SCIM group +func (g *GroupHandler) Create(r *http.Request, attributes scim.ResourceAttributes) (scim.Resource, error) { + displayName, err := getRequiredResource[string](attributes, displayNameAttr) + if err != nil { + level.Error(g.logger).Log("msg", "failed to get displayName", "err", err) + return scim.Resource{}, err + } + + // Microsoft’s SCIM implementation (Entra ID) imposes additional constraints—like enforcing uniqueness on a group’s + // displayName—that the SCIM spec itself does not mandate. + // In effect, Microsoft’s implementation diverges from strict SCIM compliance by making displayName behave like a unique key. + // SCIM only mandates that each group’s "id" is unique + _, err = g.ds.ScimGroupByDisplayName(r.Context(), displayName) + switch { + case err != nil && !fleet.IsNotFound(err): + level.Error(g.logger).Log("msg", "failed to check for displayName uniqueness", displayNameAttr, displayName, "err", err) + return scim.Resource{}, err + case err == nil: + level.Info(g.logger).Log("msg", "group already exists", displayNameAttr, displayName) + return scim.Resource{}, errors.ScimErrorUniqueness + } + + group, err := createGroupFromAttributes(attributes) + if err != nil { + level.Error(g.logger).Log("msg", "failed to create group from attributes", displayNameAttr, displayName, "err", err) + return scim.Resource{}, err + } + group.ID, err = g.ds.CreateScimGroup(r.Context(), group) + if err != nil { + return scim.Resource{}, err + } + + return createGroupResource(group), nil +} + +func createGroupFromAttributes(attributes scim.ResourceAttributes) (*fleet.ScimGroup, error) { + group := fleet.ScimGroup{} + var err error + group.DisplayName, err = getRequiredResource[string](attributes, displayNameAttr) + if err != nil { + return nil, err + } + group.ExternalID, err = getOptionalResource[string](attributes, externalIdAttr) + if err != nil { + return nil, err + } + + // Process members + members, err := getComplexResourceSlice(attributes, membersAttr) + if err != nil { + return nil, err + } + userIDs := make([]uint, 0, len(members)) + for _, member := range members { + // Get the value attribute which contains the user ID + valueIntf, ok := member["value"] + if !ok || valueIntf == nil { + continue + } + valueStr, ok := valueIntf.(string) + if !ok { + return nil, errors.ScimErrorBadParams([]string{"value"}) + } + + // Extract user ID from the value + userID, err := extractUserIDFromValue(valueStr) + if err != nil { + return nil, errors.ScimErrorBadParams([]string{"value"}) + } + userIDs = append(userIDs, userID) + } + group.ScimUsers = userIDs + + return &group, nil +} + +// Get the Scim group by ID. The group id is of the format: group-123 +// SCIM resource IDs must be unique across all resources. +func (g *GroupHandler) Get(r *http.Request, id string) (scim.Resource, error) { + idUint, err := extractGroupIDFromValue(id) + if err != nil { + level.Info(g.logger).Log("msg", "failed to parse id", "id", id, "err", err) + return scim.Resource{}, errors.ScimErrorResourceNotFound(id) + } + + group, err := g.ds.ScimGroupByID(r.Context(), idUint) + switch { + case fleet.IsNotFound(err): + level.Info(g.logger).Log("msg", "failed to find group", "id", id) + return scim.Resource{}, errors.ScimErrorResourceNotFound(id) + case err != nil: + level.Error(g.logger).Log("msg", "failed to get group", "id", id, "err", err) + return scim.Resource{}, err + } + + return createGroupResource(group), nil +} + +func createGroupResource(group *fleet.ScimGroup) scim.Resource { + groupResource := scim.Resource{} + groupResource.ID = scimGroupID(group.ID) + if group.ExternalID != nil { + groupResource.ExternalID = optional.NewString(*group.ExternalID) + } + groupResource.Attributes = scim.ResourceAttributes{} + groupResource.Attributes[displayNameAttr] = group.DisplayName + + // Add members if any + if len(group.ScimUsers) > 0 { + members := make([]scim.ResourceAttributes, 0, len(group.ScimUsers)) + for _, userID := range group.ScimUsers { + members = append(members, map[string]interface{}{ + "value": scimUserID(userID), + "$ref": "Users/" + scimUserID(userID), + "type": "User", + }) + } + groupResource.Attributes[membersAttr] = members + } + + return groupResource +} + +func (g *GroupHandler) GetAll(r *http.Request, params scim.ListRequestParams) (scim.Page, error) { + page := params.StartIndex + if page < 1 { + page = 1 + } + count := params.Count + if count > maxResults { + return scim.Page{}, errors.ScimErrorTooMany + } + if count < 1 { + count = maxResults + } + + opts := fleet.ScimListOptions{ + Page: uint(page), // nolint:gosec // ignore G115 + PerPage: uint(count), // nolint:gosec // ignore G115 + } + + resourceFilter := r.URL.Query().Get("filter") + if resourceFilter != "" { + level.Info(g.logger).Log("msg", "group filter not supported", "filter", resourceFilter) + return scim.Page{}, nil + } + + groups, totalResults, err := g.ds.ListScimGroups(r.Context(), opts) + if err != nil { + level.Error(g.logger).Log("msg", "failed to list groups", "err", err) + return scim.Page{}, err + } + + result := scim.Page{ + TotalResults: int(totalResults), // nolint:gosec // ignore G115 + Resources: make([]scim.Resource, 0, len(groups)), + } + for i := range groups { + result.Resources = append(result.Resources, createGroupResource(&groups[i])) + } + + return result, nil +} + +func (g *GroupHandler) Replace(r *http.Request, id string, attributes scim.ResourceAttributes) (scim.Resource, error) { + idUint, err := extractGroupIDFromValue(id) + if err != nil { + level.Info(g.logger).Log("msg", "failed to parse id", "id", id, "err", err) + return scim.Resource{}, errors.ScimErrorResourceNotFound(id) + } + + group, err := createGroupFromAttributes(attributes) + if err != nil { + level.Error(g.logger).Log("msg", "failed to create group from attributes", "id", id, "err", err) + return scim.Resource{}, err + } + group.ID = idUint + err = g.ds.ReplaceScimGroup(r.Context(), group) + switch { + case fleet.IsNotFound(err): + level.Info(g.logger).Log("msg", "failed to find group to replace", "id", id) + return scim.Resource{}, errors.ScimErrorResourceNotFound(id) + case err != nil: + level.Error(g.logger).Log("msg", "failed to replace group", "id", id, "err", err) + return scim.Resource{}, err + } + + return createGroupResource(group), nil +} + +func (g *GroupHandler) Delete(r *http.Request, id string) error { + idUint, err := extractGroupIDFromValue(id) + if err != nil { + level.Info(g.logger).Log("msg", "failed to parse id", "id", id, "err", err) + return errors.ScimErrorResourceNotFound(id) + } + err = g.ds.DeleteScimGroup(r.Context(), idUint) + switch { + case fleet.IsNotFound(err): + level.Info(g.logger).Log("msg", "failed to find group to delete", "id", id) + return errors.ScimErrorResourceNotFound(id) + case err != nil: + level.Error(g.logger).Log("msg", "failed to delete group", "id", id, "err", err) + return err + } + return nil +} + +// Patch +// Only supporting replacing the "displayName" attribute. +// Note: Okta does not use PATCH endpoint to update groups (2025/04/01) +func (g *GroupHandler) Patch(r *http.Request, id string, operations []scim.PatchOperation) (scim.Resource, error) { + idUint, err := extractGroupIDFromValue(id) + if err != nil { + level.Info(g.logger).Log("msg", "failed to parse id", "id", id, "err", err) + return scim.Resource{}, errors.ScimErrorResourceNotFound(id) + } + group, err := g.ds.ScimGroupByID(r.Context(), idUint) + switch { + case fleet.IsNotFound(err): + level.Info(g.logger).Log("msg", "failed to find group to patch", "id", id) + return scim.Resource{}, errors.ScimErrorResourceNotFound(id) + case err != nil: + level.Error(g.logger).Log("msg", "failed to get group to patch", "id", id, "err", err) + return scim.Resource{}, err + } + + for _, op := range operations { + if op.Op != "replace" { + level.Info(g.logger).Log("msg", "unsupported patch operation", "op", op.Op) + return scim.Resource{}, errors.ScimErrorBadParams([]string{fmt.Sprintf("%v", op)}) + } + switch { + case op.Path == nil: + newValues, ok := op.Value.(map[string]interface{}) + if !ok { + level.Info(g.logger).Log("msg", "unsupported patch value", "value", op.Value) + return scim.Resource{}, errors.ScimErrorBadParams([]string{fmt.Sprintf("%v", op)}) + } + if len(newValues) != 1 { + level.Info(g.logger).Log("msg", "too many patch values", "value", op.Value) + return scim.Resource{}, errors.ScimErrorBadParams([]string{fmt.Sprintf("%v", op)}) + } + displayName, err := getRequiredResource[string](newValues, displayNameAttr) + if err != nil { + level.Info(g.logger).Log("msg", "failed to get active value", "value", op.Value) + return scim.Resource{}, err + } + group.DisplayName = displayName + case op.Path.String() == displayNameAttr: + displayName, ok := op.Value.(string) + if !ok { + level.Error(g.logger).Log("msg", "unsupported 'displayName' patch value", "value", op.Value) + return scim.Resource{}, errors.ScimErrorBadParams([]string{fmt.Sprintf("%v", op)}) + } + group.DisplayName = displayName + default: + level.Info(g.logger).Log("msg", "unsupported patch path", "path", op.Path) + return scim.Resource{}, errors.ScimErrorBadParams([]string{fmt.Sprintf("%v", op)}) + } + } + + err = g.ds.ReplaceScimGroup(r.Context(), group) + switch { + case fleet.IsNotFound(err): + level.Info(g.logger).Log("msg", "failed to find group to patch", "id", id) + return scim.Resource{}, errors.ScimErrorResourceNotFound(id) + case err != nil: + level.Error(g.logger).Log("msg", "failed to patch group", "id", id, "err", err) + return scim.Resource{}, err + } + + return createGroupResource(group), nil +} + +func scimGroupID(groupID uint) string { + return fmt.Sprintf("group-%d", groupID) +} + +// extractGroupIDFromValue extracts the group ID from a value like "group-123" +func extractGroupIDFromValue(value string) (uint, error) { + if !strings.HasPrefix(value, "group-") { + return 0, fmt.Errorf("value %q does not match the expected format 'group-'", value) + } + + idStr := strings.TrimPrefix(value, "group-") + id, err := strconv.ParseUint(idStr, 10, 64) + if err != nil { + return 0, fmt.Errorf("failed to parse group ID from value %q: %w", value, err) + } + + return uint(id), nil +} diff --git a/ee/server/scim/scim.go b/ee/server/scim/scim.go index ac8ed0c744..b65f4cdca4 100644 --- a/ee/server/scim/scim.go +++ b/ee/server/scim/scim.go @@ -28,15 +28,13 @@ func RegisterSCIM( logger kitlog.Logger, ) error { config := scim.ServiceProviderConfig{ - // TODO: DocumentationURI and Authentication scheme DocumentationURI: optional.NewString("https://fleetdm.com/docs/get-started/why-fleet"), - SupportFiltering: true, - SupportPatch: true, MaxResults: maxResults, } // The common attributes are id, externalId, and meta. // In practice only meta.resourceType is required, while the other four (created, lastModified, location, and version) are not strictly required. + // RFC: https://tools.ietf.org/html/rfc7643#section-4.1 userSchema := schema.Schema{ ID: "urn:ietf:params:scim:schemas:core:2.0:User", Name: optional.NewString("User"), @@ -85,6 +83,63 @@ func RegisterSCIM( Description: optional.NewString("A Boolean value indicating the User's administrative status."), Name: "active", })), + schema.ComplexCoreAttribute(schema.ComplexParams{ + Description: optional.NewString("A list of groups to which the user belongs, either through direct membership, through nested groups, or dynamically calculated."), + MultiValued: true, + Mutability: schema.AttributeMutabilityReadOnly(), + Name: "groups", + SubAttributes: []schema.SimpleParams{ + schema.SimpleStringParams(schema.StringParams{ + Description: optional.NewString("The identifier of the User's group."), + Mutability: schema.AttributeMutabilityReadOnly(), + Name: "value", + }), + schema.SimpleReferenceParams(schema.ReferenceParams{ + Description: optional.NewString("The URI of the corresponding 'Group' resource to which the user belongs."), + Mutability: schema.AttributeMutabilityReadOnly(), + Name: "$ref", + ReferenceTypes: []schema.AttributeReferenceType{"Group"}, + }), + }, + }), + }, + } + + // RFC: https://tools.ietf.org/html/rfc7643#section-4.2 + groupSchema := schema.Schema{ + ID: "urn:ietf:params:scim:schemas:core:2.0:Group", + Name: optional.NewString("Group"), + Description: optional.NewString("SCIM Group"), + Attributes: []schema.CoreAttribute{ + schema.SimpleCoreAttribute(schema.SimpleStringParams(schema.StringParams{ + Description: optional.NewString("A human-readable name for the Group. REQUIRED."), + Name: "displayName", + Required: true, + })), + schema.ComplexCoreAttribute(schema.ComplexParams{ + Description: optional.NewString("A list of members of the Group."), + MultiValued: true, + Name: "members", + SubAttributes: []schema.SimpleParams{ + schema.SimpleStringParams(schema.StringParams{ + Description: optional.NewString("Identifier of the member of this Group."), + Mutability: schema.AttributeMutabilityImmutable(), + Name: "value", + }), + schema.SimpleReferenceParams(schema.ReferenceParams{ + Description: optional.NewString("The URI corresponding to a SCIM resource that is a member of this Group."), + Mutability: schema.AttributeMutabilityImmutable(), + Name: "$ref", + ReferenceTypes: []schema.AttributeReferenceType{"User"}, + }), + schema.SimpleStringParams(schema.StringParams{ + CanonicalValues: []string{"User"}, + Description: optional.NewString("A label indicating the type of resource, e.g., 'User' or 'Group'."), + Mutability: schema.AttributeMutabilityImmutable(), + Name: "type", + }), + }, + }), }, } @@ -98,6 +153,14 @@ func RegisterSCIM( Schema: userSchema, Handler: NewUserHandler(ds, scimLogger), }, + { + ID: optional.NewString("Group"), + Name: "Group", + Endpoint: "/Groups", + Description: optional.NewString("Group"), + Schema: groupSchema, + Handler: NewGroupHandler(ds, scimLogger), + }, } serverArgs := &scim.ServerArgs{ @@ -132,6 +195,8 @@ func RegisterSCIM( return handler } + // We cannot use Go URL path pattern like {version} because the http.StripPrefix method + // that gets us to the root SCIM path does not support wildcards: https://github.com/golang/go/issues/64909 mux.Handle("/api/v1/fleet/scim/", applyMiddleware("/api/v1/fleet/scim", server)) mux.Handle("/api/latest/fleet/scim/", applyMiddleware("/api/latest/fleet/scim", server)) return nil diff --git a/ee/server/scim/users.go b/ee/server/scim/users.go index cbed010fdf..beef4545a4 100644 --- a/ee/server/scim/users.go +++ b/ee/server/scim/users.go @@ -28,6 +28,7 @@ const ( familyNameAttr = "familyName" activeAttr = "active" emailsAttr = "emails" + groupsAttr = "groups" ) type UserHandler struct { @@ -50,7 +51,11 @@ func (u *UserHandler) Create(r *http.Request, attributes scim.ResourceAttributes return scim.Resource{}, err } _, err = u.ds.ScimUserByUserName(r.Context(), userName) - if !fleet.IsNotFound(err) { + switch { + case err != nil && !fleet.IsNotFound(err): + level.Error(u.logger).Log("msg", "failed to check for userName uniqueness", userNameAttr, userName, "err", err) + return scim.Resource{}, err + case err == nil: level.Info(u.logger).Log("msg", "user already exists", userNameAttr, userName) return scim.Resource{}, errors.ScimErrorUniqueness } @@ -187,13 +192,13 @@ func getComplexResourceSlice(attributes scim.ResourceAttributes, key string) ([] } func (u *UserHandler) Get(r *http.Request, id string) (scim.Resource, error) { - idUint, err := strconv.ParseUint(id, 10, 64) + idUint, err := extractUserIDFromValue(id) if err != nil { level.Info(u.logger).Log("msg", "failed to parse id", "id", id, "err", err) return scim.Resource{}, errors.ScimErrorResourceNotFound(id) } - user, err := u.ds.ScimUserByID(r.Context(), uint(idUint)) + user, err := u.ds.ScimUserByID(r.Context(), idUint) switch { case fleet.IsNotFound(err): level.Info(u.logger).Log("msg", "failed to find user", "id", id) @@ -208,7 +213,7 @@ func (u *UserHandler) Get(r *http.Request, id string) (scim.Resource, error) { func createUserResource(user *fleet.ScimUser) scim.Resource { userResource := scim.Resource{} - userResource.ID = fmt.Sprintf("%d", user.ID) + userResource.ID = scimUserID(user.ID) if user.ExternalID != nil { userResource.ExternalID = optional.NewString(*user.ExternalID) } @@ -241,6 +246,16 @@ func createUserResource(user *fleet.ScimUser) scim.Resource { } userResource.Attributes[emailsAttr] = emails } + if len(user.Groups) > 0 { + groups := make([]scim.ResourceAttributes, 0, len(user.Groups)) + for _, groupID := range user.Groups { + groups = append(groups, map[string]interface{}{ + "value": scimGroupID(groupID), + "$ref": "Groups/" + scimGroupID(groupID), + }) + } + userResource.Attributes[groupsAttr] = groups + } return userResource } @@ -272,8 +287,10 @@ func (u *UserHandler) GetAll(r *http.Request, params scim.ListRequestParams) (sc } opts := fleet.ScimUsersListOptions{ - Page: uint(page), // nolint:gosec // ignore G115 - PerPage: uint(count), // nolint:gosec // ignore G115 + ScimListOptions: fleet.ScimListOptions{ + Page: uint(page), // nolint:gosec // ignore G115 + PerPage: uint(count), // nolint:gosec // ignore G115 + }, } resourceFilter := r.URL.Query().Get("filter") if resourceFilter != "" { @@ -318,7 +335,7 @@ func (u *UserHandler) GetAll(r *http.Request, params scim.ListRequestParams) (sc } func (u *UserHandler) Replace(r *http.Request, id string, attributes scim.ResourceAttributes) (scim.Resource, error) { - idUint, err := strconv.ParseUint(id, 10, 64) + idUint, err := extractUserIDFromValue(id) if err != nil { level.Info(u.logger).Log("msg", "failed to parse id", "id", id, "err", err) return scim.Resource{}, errors.ScimErrorResourceNotFound(id) @@ -329,7 +346,7 @@ func (u *UserHandler) Replace(r *http.Request, id string, attributes scim.Resour level.Error(u.logger).Log("msg", "failed to create user from attributes", "id", id, "err", err) return scim.Resource{}, err } - user.ID = uint(idUint) + user.ID = idUint err = u.ds.ReplaceScimUser(r.Context(), user) switch { case fleet.IsNotFound(err): @@ -347,12 +364,12 @@ func (u *UserHandler) Replace(r *http.Request, id string, attributes scim.Resour // https://datatracker.ietf.org/doc/html/rfc7644#section-3.6 // MUST return a 404 (Not Found) error code for all operations associated with the previously deleted resource func (u *UserHandler) Delete(r *http.Request, id string) error { - idUint, err := strconv.ParseUint(id, 10, 64) + idUint, err := extractUserIDFromValue(id) if err != nil { level.Info(u.logger).Log("msg", "failed to parse id", "id", id, "err", err) return errors.ScimErrorResourceNotFound(id) } - err = u.ds.DeleteScimUser(r.Context(), uint(idUint)) + err = u.ds.DeleteScimUser(r.Context(), idUint) switch { case fleet.IsNotFound(err): level.Info(u.logger).Log("msg", "failed to find user to delete", "id", id) @@ -368,12 +385,12 @@ func (u *UserHandler) Delete(r *http.Request, id string) error { // Okta only requires patching the "active" attribute: // https://developer.okta.com/docs/api/openapi/okta-scim/guides/scim-20/#update-a-specific-user-patch func (u *UserHandler) Patch(r *http.Request, id string, operations []scim.PatchOperation) (scim.Resource, error) { - idUint, err := strconv.ParseUint(id, 10, 64) + idUint, err := extractUserIDFromValue(id) if err != nil { level.Info(u.logger).Log("msg", "failed to parse id", "id", id, "err", err) return scim.Resource{}, errors.ScimErrorResourceNotFound(id) } - user, err := u.ds.ScimUserByID(r.Context(), uint(idUint)) + user, err := u.ds.ScimUserByID(r.Context(), idUint) switch { case fleet.IsNotFound(err): level.Info(u.logger).Log("msg", "failed to find user to patch", "id", id) @@ -453,3 +470,17 @@ func removeWhitespace(str string) string { return r }, str) } + +func scimUserID(userID uint) string { + return fmt.Sprintf("%d", userID) +} + +// extractUserIDFromValue extracts the user ID from a value like "123" +func extractUserIDFromValue(value string) (uint, error) { + id, err := strconv.ParseUint(value, 10, 64) + if err != nil { + return 0, err + } + + return uint(id), nil +} diff --git a/server/datastore/mysql/scim.go b/server/datastore/mysql/scim.go index 8b29e01164..b521a61180 100644 --- a/server/datastore/mysql/scim.go +++ b/server/datastore/mysql/scim.go @@ -4,15 +4,28 @@ import ( "context" "database/sql" "errors" + "fmt" "strings" "github.com/fleetdm/fleet/v4/server/contexts/ctxerr" + "github.com/fleetdm/fleet/v4/server/datastore/mysql/common_mysql" "github.com/fleetdm/fleet/v4/server/fleet" "github.com/jmoiron/sqlx" ) +const ( + // SCIMMaxFieldLength is the maximum length for SCIM user fields + SCIMMaxFieldLength = 255 + + SCIMDefaultResourcesPerPage = 100 +) + // CreateScimUser creates a new SCIM user in the database func (ds *Datastore) CreateScimUser(ctx context.Context, user *fleet.ScimUser) (uint, error) { + if err := validateScimUserFields(user); err != nil { + return 0, err + } + var userID uint err := ds.withRetryTxx(ctx, func(tx sqlx.ExtContext) error { const insertUserQuery = ` @@ -68,6 +81,13 @@ func (ds *Datastore) ScimUserByID(ctx context.Context, id uint) (*fleet.ScimUser } user.Emails = emails + // Get the user's groups + groups, err := ds.getScimUserGroups(ctx, id) + if err != nil { + return nil, err + } + user.Groups = groups + return user, nil } @@ -95,11 +115,22 @@ func (ds *Datastore) ScimUserByUserName(ctx context.Context, userName string) (* } user.Emails = emails + // Get the user's groups + groups, err := ds.getScimUserGroups(ctx, user.ID) + if err != nil { + return nil, err + } + user.Groups = groups + return user, nil } // ReplaceScimUser replaces an existing SCIM user in the database func (ds *Datastore) ReplaceScimUser(ctx context.Context, user *fleet.ScimUser) error { + if err := validateScimUserFields(user); err != nil { + return err + } + return ds.withRetryTxx(ctx, func(tx sqlx.ExtContext) error { // Update the SCIM user const updateUserQuery = ` @@ -144,7 +175,19 @@ func (ds *Datastore) ReplaceScimUser(ctx context.Context, user *fleet.ScimUser) return ctxerr.Wrap(ctx, err, "delete scim user emails") } - return insertEmails(ctx, tx, user) + err = insertEmails(ctx, tx, user) + if err != nil { + return err + } + + // Get the user's groups + groups, err := ds.getScimUserGroups(ctx, user.ID) + if err != nil { + return err + } + user.Groups = groups + + return nil }) } @@ -218,7 +261,7 @@ func (ds *Datastore) ListScimUsers(ctx context.Context, opts fleet.ScimUsersList opts.Page = 1 } if opts.PerPage == 0 { - opts.PerPage = 100 + opts.PerPage = SCIMDefaultResourcesPerPage } // Calculate offset for pagination @@ -308,6 +351,37 @@ func (ds *Datastore) ListScimUsers(ctx context.Context, opts fleet.ScimUsersList } } + // Fetch groups for all users in a single query + groupQuery, groupArgs, err := sqlx.In(` + SELECT + scim_user_id, group_id + FROM scim_user_group + WHERE scim_user_id IN (?) + ORDER BY group_id ASC + `, userIDs) + if err != nil { + return nil, 0, ctxerr.Wrap(ctx, err, "prepare groups query") + } + + // Execute the group query + type userGroup struct { + UserID uint `db:"scim_user_id"` + GroupID uint `db:"group_id"` + } + var allUserGroups []userGroup + if err := sqlx.SelectContext(ctx, ds.reader(ctx), &allUserGroups, groupQuery, groupArgs...); err != nil { + if !errors.Is(err, sql.ErrNoRows) { + return nil, 0, ctxerr.Wrap(ctx, err, "select scim user groups") + } + } + + // Associate groups with their users + for _, ug := range allUserGroups { + if user, ok := userMap[ug.UserID]; ok { + user.Groups = append(user.Groups, ug.GroupID) + } + } + return users, totalResults, nil } @@ -329,3 +403,401 @@ func (ds *Datastore) getScimUserEmails(ctx context.Context, userID uint) ([]flee } return emails, nil } + +// getScimUserGroups retrieves all group IDs for a SCIM user +func (ds *Datastore) getScimUserGroups(ctx context.Context, userID uint) ([]uint, error) { + const query = ` + SELECT + group_id + FROM scim_user_group + WHERE scim_user_id = ? ORDER BY group_id ASC + ` + var groupIDs []uint + err := sqlx.SelectContext(ctx, ds.reader(ctx), &groupIDs, query, userID) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, ctxerr.Wrap(ctx, err, "select scim user groups") + } + return groupIDs, nil +} + +// validateScimUserFields checks if the user fields exceed the maximum allowed length +func validateScimUserFields(user *fleet.ScimUser) error { + if user.ExternalID != nil && len(*user.ExternalID) > SCIMMaxFieldLength { + return fmt.Errorf("external_id exceeds maximum length of %d characters", SCIMMaxFieldLength) + } + if len(user.UserName) > SCIMMaxFieldLength { + return fmt.Errorf("user_name exceeds maximum length of %d characters", SCIMMaxFieldLength) + } + if user.GivenName != nil && len(*user.GivenName) > SCIMMaxFieldLength { + return fmt.Errorf("given_name exceeds maximum length of %d characters", SCIMMaxFieldLength) + } + if user.FamilyName != nil && len(*user.FamilyName) > SCIMMaxFieldLength { + return fmt.Errorf("family_name exceeds maximum length of %d characters", SCIMMaxFieldLength) + } + return nil +} + +// validateScimGroupFields checks if the group fields exceed the maximum allowed length +func validateScimGroupFields(group *fleet.ScimGroup) error { + if group.ExternalID != nil && len(*group.ExternalID) > SCIMMaxFieldLength { + return fmt.Errorf("external_id exceeds maximum length of %d characters", SCIMMaxFieldLength) + } + if len(group.DisplayName) > SCIMMaxFieldLength { + return fmt.Errorf("display_name exceeds maximum length of %d characters", SCIMMaxFieldLength) + } + return nil +} + +// CreateScimGroup creates a new SCIM group in the database +func (ds *Datastore) CreateScimGroup(ctx context.Context, group *fleet.ScimGroup) (uint, error) { + if err := validateScimGroupFields(group); err != nil { + return 0, err + } + + var groupID uint + err := ds.withRetryTxx(ctx, func(tx sqlx.ExtContext) error { + const insertGroupQuery = ` + INSERT INTO scim_groups ( + external_id, display_name + ) VALUES (?, ?)` + result, err := tx.ExecContext( + ctx, + insertGroupQuery, + group.ExternalID, + group.DisplayName, + ) + if err != nil { + return ctxerr.Wrap(ctx, err, "insert scim group") + } + + id, err := result.LastInsertId() + if err != nil { + return ctxerr.Wrap(ctx, err, "insert scim group last insert id") + } + group.ID = uint(id) // nolint:gosec // dismiss G115 + groupID = group.ID + + // Insert user-group relationships if any + if len(group.ScimUsers) > 0 { + return insertScimGroupUsers(ctx, tx, group.ID, group.ScimUsers) + } + + return nil + }) + return groupID, err +} + +// insertScimGroupUsers inserts the relationships between a SCIM group and its users +func insertScimGroupUsers(ctx context.Context, tx sqlx.ExtContext, groupID uint, userIDs []uint) error { + if len(userIDs) == 0 { + return nil + } + + batchSize := 10000 + return common_mysql.BatchProcessSimple(userIDs, batchSize, func(userIDsInBatch []uint) error { + // Build the batch insert query + valueStrings := make([]string, 0, len(userIDsInBatch)) + valueArgs := make([]interface{}, 0, len(userIDsInBatch)*2) + + for _, userID := range userIDsInBatch { + valueStrings = append(valueStrings, "(?, ?)") + valueArgs = append(valueArgs, userID, groupID) + } + + // Construct the batch insert query + insertQuery := ` + INSERT INTO scim_user_group ( + scim_user_id, group_id + ) VALUES ` + strings.Join(valueStrings, ",") + + // Execute the batch insert + _, err := tx.ExecContext(ctx, insertQuery, valueArgs...) + if err != nil { + return ctxerr.Wrap(ctx, err, "batch insert scim group users") + } + return nil + }) +} + +// ScimGroupByID retrieves a SCIM group by ID +func (ds *Datastore) ScimGroupByID(ctx context.Context, id uint) (*fleet.ScimGroup, error) { + const query = ` + SELECT + id, external_id, display_name + FROM scim_groups + WHERE id = ? + ` + group := &fleet.ScimGroup{} + err := sqlx.GetContext(ctx, ds.reader(ctx), group, query, id) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, notFound("scim group").WithID(id) + } + return nil, ctxerr.Wrap(ctx, err, "select scim group") + } + + // Get the group's users + users, err := ds.getScimGroupUsers(ctx, ds.reader(ctx), id) + if err != nil { + return nil, err + } + group.ScimUsers = users + + return group, nil +} + +// ScimGroupByDisplayName retrieves a SCIM group by display name +func (ds *Datastore) ScimGroupByDisplayName(ctx context.Context, displayName string) (*fleet.ScimGroup, error) { + const query = ` + SELECT + id, external_id, display_name + FROM scim_groups + WHERE display_name = ? + ` + group := &fleet.ScimGroup{} + err := sqlx.GetContext(ctx, ds.reader(ctx), group, query, displayName) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, notFound("scim group") + } + return nil, ctxerr.Wrap(ctx, err, "select scim group by displayName") + } + + // Get the group's users + users, err := ds.getScimGroupUsers(ctx, ds.reader(ctx), group.ID) + if err != nil { + return nil, err + } + group.ScimUsers = users + + return group, nil +} + +// getScimGroupUsers retrieves all user IDs for a SCIM group +func (ds *Datastore) getScimGroupUsers(ctx context.Context, q sqlx.QueryerContext, groupID uint) ([]uint, error) { + const query = ` + SELECT + scim_user_id + FROM scim_user_group + WHERE group_id = ? ORDER BY scim_user_id ASC + ` + var userIDs []uint + err := sqlx.SelectContext(ctx, q, &userIDs, query, groupID) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return nil, ctxerr.Wrap(ctx, err, "select scim group users") + } + return userIDs, nil +} + +// ReplaceScimGroup replaces an existing SCIM group in the database +func (ds *Datastore) ReplaceScimGroup(ctx context.Context, group *fleet.ScimGroup) error { + if err := validateScimGroupFields(group); err != nil { + return err + } + + return ds.withRetryTxx(ctx, func(tx sqlx.ExtContext) error { + // Update the SCIM group + const updateGroupQuery = ` + UPDATE scim_groups SET + external_id = ?, + display_name = ? + WHERE id = ?` + result, err := tx.ExecContext( + ctx, + updateGroupQuery, + group.ExternalID, + group.DisplayName, + group.ID, + ) + if err != nil { + return ctxerr.Wrap(ctx, err, "update scim group") + } + + rowsAffected, err := result.RowsAffected() + if err != nil { + return ctxerr.Wrap(ctx, err, "get rows affected for update scim group") + } + if rowsAffected == 0 { + return notFound("scim group").WithID(group.ID) + } + + // Get existing user-group relationships + existingUsers, err := ds.getScimGroupUsers(ctx, tx, group.ID) + if err != nil { + return ctxerr.Wrap(ctx, err, "get existing scim group users") + } + + // Create maps for efficient lookup + existingUserMap := make(map[uint]bool) + for _, userID := range existingUsers { + existingUserMap[userID] = true + } + + newUserMap := make(map[uint]bool) + for _, userID := range group.ScimUsers { + newUserMap[userID] = true + } + + // Find users to add (in new but not in existing) + var usersToAdd []uint + for _, userID := range group.ScimUsers { + if !existingUserMap[userID] { + usersToAdd = append(usersToAdd, userID) + } + } + + // Find users to remove (in existing but not in new) + var usersToRemove []uint + for _, userID := range existingUsers { + if !newUserMap[userID] { + usersToRemove = append(usersToRemove, userID) + } + } + + // Add new user-group relationships + if len(usersToAdd) > 0 { + err = insertScimGroupUsers(ctx, tx, group.ID, usersToAdd) + if err != nil { + return ctxerr.Wrap(ctx, err, "insert new scim group users") + } + } + + // Remove old user-group relationships + if len(usersToRemove) > 0 { + batchSize := 10000 + return common_mysql.BatchProcessSimple(usersToRemove, batchSize, func(usersToRemoveInBatch []uint) error { + params := make([]interface{}, len(usersToRemoveInBatch)+1) + params[0] = group.ID + for i, userID := range usersToRemoveInBatch { + params[i+1] = userID + } + + deleteQuery := "DELETE FROM scim_user_group WHERE group_id = ? AND scim_user_id IN (" + + strings.Repeat("?, ", len(usersToRemoveInBatch)-1) + "?)" + + _, err = tx.ExecContext(ctx, deleteQuery, params...) + if err != nil { + return ctxerr.Wrap(ctx, err, "delete removed scim group users") + } + return nil + }) + } + + return nil + }) +} + +// DeleteScimGroup deletes a SCIM group from the database +func (ds *Datastore) DeleteScimGroup(ctx context.Context, id uint) error { + return ds.withRetryTxx(ctx, func(tx sqlx.ExtContext) error { + // Delete the group + const deleteGroupQuery = `DELETE FROM scim_groups WHERE id = ?` + result, err := tx.ExecContext(ctx, deleteGroupQuery, id) + if err != nil { + return ctxerr.Wrap(ctx, err, "delete scim group") + } + + // Check if the group existed + rowsAffected, err := result.RowsAffected() + if err != nil { + return ctxerr.Wrap(ctx, err, "get rows affected for delete scim group") + } + if rowsAffected == 0 { + return notFound("scim group").WithID(id) + } + + return nil + }) +} + +// ListScimGroups retrieves a list of SCIM groups with pagination +func (ds *Datastore) ListScimGroups(ctx context.Context, opts fleet.ScimListOptions) (groups []fleet.ScimGroup, totalResults uint, err error) { + // Default pagination values if not provided + if opts.Page == 0 { + opts.Page = 1 + } + if opts.PerPage == 0 { + opts.PerPage = SCIMDefaultResourcesPerPage + } + + // Calculate offset for pagination + offset := (opts.Page - 1) * opts.PerPage + + // Build the query + baseQuery := ` + SELECT DISTINCT + scim_groups.id, external_id, display_name + FROM scim_groups + ` + + // First, get the total count without pagination + countQuery := "SELECT COUNT(DISTINCT id) FROM (" + baseQuery + ") AS filtered_groups" + err = sqlx.GetContext(ctx, ds.reader(ctx), &totalResults, countQuery) + if err != nil { + return nil, 0, ctxerr.Wrap(ctx, err, "count total scim groups") + } + + // Add pagination to the main query + query := baseQuery + " ORDER BY scim_groups.id LIMIT ? OFFSET ?" + params := []interface{}{opts.PerPage, offset} + + // Execute the query + err = sqlx.SelectContext(ctx, ds.reader(ctx), &groups, query, params...) + if err != nil { + return nil, 0, ctxerr.Wrap(ctx, err, "list scim groups") + } + + // Process the results + groupIDs := make([]uint, 0, len(groups)) + groupMap := make(map[uint]*fleet.ScimGroup, len(groups)) + + for i, group := range groups { + groupIDs = append(groupIDs, group.ID) + groupMap[group.ID] = &groups[i] + groups[i].ScimUsers = []uint{} // Initialize empty user list for each group + } + + // If no groups found, return empty slice + if len(groups) == 0 { + return groups, totalResults, nil + } + + // Fetch users for all groups in a single query + userQuery, args, err := sqlx.In(` + SELECT + group_id, scim_user_id + FROM scim_user_group + WHERE group_id IN (?) + ORDER BY scim_user_id ASC + `, groupIDs) + if err != nil { + return nil, 0, ctxerr.Wrap(ctx, err, "prepare users query") + } + + // Execute the user query + type groupUser struct { + GroupID uint `db:"group_id"` + UserID uint `db:"scim_user_id"` + } + var allGroupUsers []groupUser + if err := sqlx.SelectContext(ctx, ds.reader(ctx), &allGroupUsers, userQuery, args...); err != nil { + if !errors.Is(err, sql.ErrNoRows) { + return nil, 0, ctxerr.Wrap(ctx, err, "select scim group users") + } + } + + // Associate users with their groups + for _, gu := range allGroupUsers { + if group, ok := groupMap[gu.GroupID]; ok { + group.ScimUsers = append(group.ScimUsers, gu.UserID) + } + } + + return groups, totalResults, nil +} diff --git a/server/datastore/mysql/scim_test.go b/server/datastore/mysql/scim_test.go index e07eca156c..914b657230 100644 --- a/server/datastore/mysql/scim_test.go +++ b/server/datastore/mysql/scim_test.go @@ -2,6 +2,8 @@ package mysql import ( "context" + "sort" + "strings" "testing" "github.com/fleetdm/fleet/v4/server/fleet" @@ -18,15 +20,25 @@ func TestScim(t *testing.T) { fn func(t *testing.T, ds *Datastore) }{ {"ScimUserCreate", testScimUserCreate}, + {"ScimUserCreateValidation", testScimUserCreateValidation}, {"ScimUserByID", testScimUserByID}, {"ScimUserByUserName", testScimUserByUserName}, {"ReplaceScimUser", testReplaceScimUser}, + {"ReplaceScimUserValidation", testScimUserReplaceValidation}, {"DeleteScimUser", testDeleteScimUser}, {"ListScimUsers", testListScimUsers}, + {"ScimGroupCreate", testScimGroupCreate}, + {"ScimGroupCreateValidation", testScimGroupCreateValidation}, + {"ScimGroupByID", testScimGroupByID}, + {"ScimGroupByDisplayName", testScimGroupByDisplayName}, + {"ReplaceScimGroup", testReplaceScimGroup}, + {"ReplaceScimGroupValidation", testScimGroupReplaceValidation}, + {"DeleteScimGroup", testDeleteScimGroup}, + {"ListScimGroups", testListScimGroups}, } for _, c := range cases { t.Run(c.name, func(t *testing.T) { - defer TruncateTables(t, ds, "scim_users", "scim_user_emails") + defer TruncateTables(t, ds, "scim_users", "scim_user_emails", "scim_groups", "scim_user_group") c.fn(t, ds) }) } @@ -106,6 +118,10 @@ func testScimUserCreate(t *testing.T, ds *Datastore) { func testScimUserByID(t *testing.T, ds *Datastore) { users := createTestScimUsers(t, ds) + + // Create test groups and associate them with users + groups := createTestScimGroups(t, ds, []uint{users[0].ID, users[1].ID}) + for _, tt := range users { returned, err := ds.ScimUserByID(context.Background(), tt.ID) assert.Nil(t, err) @@ -124,6 +140,39 @@ func testScimUserByID(t *testing.T, ds *Datastore) { assert.Equal(t, email.Type, returned.Emails[i].Type) assert.Equal(t, tt.ID, returned.Emails[i].ScimUserID) } + + // Verify groups + // User 0 and 1 should be in groups, User 2 should not be in any group + if tt.ID == users[0].ID || tt.ID == users[1].ID { + assert.NotEmpty(t, returned.Groups, "User should have groups") + + // Check if the user is in the expected groups + var foundInGroups bool + for _, group := range groups { + for _, userID := range group.ScimUsers { + if userID == tt.ID { + foundInGroups = true + + // Verify the group ID is in the user's Groups field + var foundGroupID bool + for _, groupID := range returned.Groups { + if groupID == group.ID { + foundGroupID = true + break + } + } + assert.True(t, foundGroupID, "User's Groups field should contain the group ID") + break + } + } + if foundInGroups { + break + } + } + assert.True(t, foundInGroups, "User should be found in at least one group") + } else { + assert.Empty(t, returned.Groups, "User should not have any groups") + } } // test missing user @@ -133,6 +182,10 @@ func testScimUserByID(t *testing.T, ds *Datastore) { func testScimUserByUserName(t *testing.T, ds *Datastore) { users := createTestScimUsers(t, ds) + + // Create test groups and associate them with users + groups := createTestScimGroups(t, ds, []uint{users[0].ID, users[1].ID}) + for _, tt := range users { returned, err := ds.ScimUserByUserName(context.Background(), tt.UserName) assert.Nil(t, err) @@ -151,6 +204,39 @@ func testScimUserByUserName(t *testing.T, ds *Datastore) { assert.Equal(t, email.Type, returned.Emails[i].Type) assert.Equal(t, tt.ID, returned.Emails[i].ScimUserID) } + + // Verify groups + // User 0 and 1 should be in groups, User 2 should not be in any group + if tt.ID == users[0].ID || tt.ID == users[1].ID { + assert.NotEmpty(t, returned.Groups, "User should have groups") + + // Check if the user is in the expected groups + var foundInGroups bool + for _, group := range groups { + for _, userID := range group.ScimUsers { + if userID == tt.ID { + foundInGroups = true + + // Verify the group ID is in the user's Groups field + var foundGroupID bool + for _, groupID := range returned.Groups { + if groupID == group.ID { + foundGroupID = true + break + } + } + assert.True(t, foundGroupID, "User's Groups field should contain the group ID") + break + } + } + if foundInGroups { + break + } + } + assert.True(t, foundInGroups, "User should be found in at least one group") + } else { + assert.Empty(t, returned.Groups, "User should not have any groups") + } } // test missing user @@ -226,7 +312,16 @@ func testReplaceScimUser(t *testing.T, ds *Datastore) { user.ID, err = ds.CreateScimUser(context.Background(), &user) require.Nil(t, err) - // Verify the user was created correctly + // Create a test group and associate it with the user + group := fleet.ScimGroup{ + DisplayName: "Test Group for User", + ExternalID: ptr.String("ext-group-for-user"), + ScimUsers: []uint{user.ID}, + } + group.ID, err = ds.CreateScimGroup(context.Background(), &group) + require.Nil(t, err) + + // Verify the user was created correctly and has the group createdUser, err := ds.ScimUserByID(context.Background(), user.ID) require.Nil(t, err) assert.Equal(t, user.UserName, createdUser.UserName) @@ -237,7 +332,11 @@ func testReplaceScimUser(t *testing.T, ds *Datastore) { assert.Equal(t, 1, len(createdUser.Emails)) assert.Equal(t, "original.user@example.com", createdUser.Emails[0].Email) - // Modify the user + // Verify the user has the group + require.Len(t, createdUser.Groups, 1) + assert.Equal(t, group.ID, createdUser.Groups[0]) + + // Modify the user and attempt to modify the Groups field updatedUser := fleet.ScimUser{ ID: user.ID, UserName: "replace-test-user", // Same username @@ -257,6 +356,7 @@ func testReplaceScimUser(t *testing.T, ds *Datastore) { Type: ptr.String("home"), }, }, + Groups: []uint{999}, // Attempt to modify Groups (should be ignored) } // Replace the user @@ -277,6 +377,25 @@ func testReplaceScimUser(t *testing.T, ds *Datastore) { assert.Equal(t, "personal.user@example.com", replacedUser.Emails[0].Email) // Alphabetical order assert.Equal(t, "updated.user@example.com", replacedUser.Emails[1].Email) + // Verify that the Groups field was NOT modified (it should still contain the original group) + require.Len(t, replacedUser.Groups, 1, "Groups field should not be modified by ReplaceScimUser") + assert.Equal(t, group.ID, replacedUser.Groups[0], "Groups field should still contain the original group") + + // Now remove the user from the group using the group methods + updatedGroup := fleet.ScimGroup{ + ID: group.ID, + DisplayName: group.DisplayName, + ExternalID: group.ExternalID, + ScimUsers: []uint{}, // Remove the user + } + err = ds.ReplaceScimGroup(context.Background(), &updatedGroup) + require.Nil(t, err) + + // Verify that the user no longer has the group + userAfterGroupUpdate, err := ds.ScimUserByID(context.Background(), user.ID) + require.Nil(t, err) + assert.Empty(t, userAfterGroupUpdate.Groups, "User should no longer have any groups") + // Test replacing a non-existent user nonExistentUser := fleet.ScimUser{ ID: 99999, // Non-existent ID @@ -389,10 +508,22 @@ func testListScimUsers(t *testing.T, ds *Datastore) { require.Nil(t, err) } + // Create a group and associate it with the first user + group := fleet.ScimGroup{ + DisplayName: "Test Group for ListUsers", + ExternalID: ptr.String("ext-group-for-list"), + ScimUsers: []uint{users[0].ID}, + } + var err error + group.ID, err = ds.CreateScimGroup(context.Background(), &group) + require.Nil(t, err) + // Test 1: List all users without filters allUsers, totalResults, err := ds.ListScimUsers(context.Background(), fleet.ScimUsersListOptions{ - Page: 1, - PerPage: 10, + ScimListOptions: fleet.ScimListOptions{ + Page: 1, + PerPage: 10, + }, }) require.Nil(t, err) assert.Equal(t, 3, len(allUsers)) @@ -404,6 +535,15 @@ func testListScimUsers(t *testing.T, ds *Datastore) { for _, testUser := range users { if u.ID == testUser.ID { foundUsers++ + + // Verify Groups field for the first user + if testUser.ID == users[0].ID { + require.Len(t, u.Groups, 1, "First user should have exactly one group") + assert.Equal(t, group.ID, u.Groups[0], "First user should be in the test group") + } else { + assert.Empty(t, u.Groups, "Other users should not have groups") + } + break } } @@ -412,8 +552,10 @@ func testListScimUsers(t *testing.T, ds *Datastore) { // Test 2: Pagination - first page with 2 items page1Users, totalPage1, err := ds.ListScimUsers(context.Background(), fleet.ScimUsersListOptions{ - Page: 1, - PerPage: 2, + ScimListOptions: fleet.ScimListOptions{ + Page: 1, + PerPage: 2, + }, }) require.Nil(t, err) assert.Equal(t, 2, len(page1Users)) @@ -421,8 +563,10 @@ func testListScimUsers(t *testing.T, ds *Datastore) { // Test 3: Pagination - second page with 2 items page2Users, totalPage2, err := ds.ListScimUsers(context.Background(), fleet.ScimUsersListOptions{ - Page: 2, - PerPage: 2, + ScimListOptions: fleet.ScimListOptions{ + Page: 2, + PerPage: 2, + }, }) require.Nil(t, err) assert.Equal(t, 1, len(page2Users)) @@ -437,8 +581,10 @@ func testListScimUsers(t *testing.T, ds *Datastore) { // Test 4: Filter by username listUsers, totalListUsers, err := ds.ListScimUsers(context.Background(), fleet.ScimUsersListOptions{ - Page: 1, - PerPage: 10, + ScimListOptions: fleet.ScimListOptions{ + Page: 1, + PerPage: 10, + }, UserNameFilter: ptr.String("list-test-user2"), }) @@ -449,8 +595,10 @@ func testListScimUsers(t *testing.T, ds *Datastore) { // Test 5: Filter by email type and value homeEmailUsers, totalHomeEmailUsers, err := ds.ListScimUsers(context.Background(), fleet.ScimUsersListOptions{ - Page: 1, - PerPage: 10, + ScimListOptions: fleet.ScimListOptions{ + Page: 1, + PerPage: 10, + }, EmailTypeFilter: ptr.String("home"), EmailValueFilter: ptr.String("personal.user2@example.com"), }) @@ -462,8 +610,10 @@ func testListScimUsers(t *testing.T, ds *Datastore) { // Test 6: Filter by email type and value - work emails workEmailUsers, totalWorkEmailUsers, err := ds.ListScimUsers(context.Background(), fleet.ScimUsersListOptions{ - Page: 1, - PerPage: 10, + ScimListOptions: fleet.ScimListOptions{ + Page: 1, + PerPage: 10, + }, EmailTypeFilter: ptr.String("work"), EmailValueFilter: ptr.String("different.user3@example.com"), }) @@ -473,8 +623,10 @@ func testListScimUsers(t *testing.T, ds *Datastore) { // Test 7: No results for non-matching filters noUsers, totalNoUsers1, err := ds.ListScimUsers(context.Background(), fleet.ScimUsersListOptions{ - Page: 1, - PerPage: 10, + ScimListOptions: fleet.ScimListOptions{ + Page: 1, + PerPage: 10, + }, UserNameFilter: ptr.String("nonexistent"), }) require.Nil(t, err) @@ -482,8 +634,10 @@ func testListScimUsers(t *testing.T, ds *Datastore) { assert.Equal(t, uint(0), totalNoUsers1) noUsers, totalNoUsers2, err := ds.ListScimUsers(context.Background(), fleet.ScimUsersListOptions{ - Page: 1, - PerPage: 10, + ScimListOptions: fleet.ScimListOptions{ + Page: 1, + PerPage: 10, + }, EmailTypeFilter: ptr.String("nonexistent"), EmailValueFilter: ptr.String("nonexistent"), }) @@ -491,3 +645,578 @@ func testListScimUsers(t *testing.T, ds *Datastore) { assert.Empty(t, noUsers) assert.Equal(t, uint(0), totalNoUsers2) } + +func testScimGroupCreate(t *testing.T, ds *Datastore) { + // Create test users first + users := createTestScimUsers(t, ds) + userIDs := make([]uint, len(users)) + for i, user := range users { + userIDs[i] = user.ID + } + + groupsToCreate := []fleet.ScimGroup{ + { + DisplayName: "Group1", + ExternalID: nil, + ScimUsers: []uint{}, + }, + { + DisplayName: "Group2", + ExternalID: ptr.String("ext-group-123"), + ScimUsers: []uint{userIDs[0]}, + }, + { + DisplayName: "Group3", + ExternalID: ptr.String("ext-group-456"), + ScimUsers: userIDs, + }, + } + + for _, g := range groupsToCreate { + var err error + groupCopy := g + groupCopy.ID, err = ds.CreateScimGroup(context.Background(), &g) + assert.Nil(t, err) + + verify, err := ds.ScimGroupByID(context.Background(), g.ID) + assert.Nil(t, err) + + assert.Equal(t, groupCopy.ID, verify.ID) + assert.Equal(t, groupCopy.DisplayName, verify.DisplayName) + assert.Equal(t, groupCopy.ExternalID, verify.ExternalID) + + // Verify users + assert.Equal(t, len(groupCopy.ScimUsers), len(verify.ScimUsers)) + if len(groupCopy.ScimUsers) > 0 { + // Sort the user IDs for comparison + sort.Slice(groupCopy.ScimUsers, func(i, j int) bool { + return groupCopy.ScimUsers[i] < groupCopy.ScimUsers[j] + }) + sort.Slice(verify.ScimUsers, func(i, j int) bool { + return verify.ScimUsers[i] < verify.ScimUsers[j] + }) + assert.Equal(t, groupCopy.ScimUsers, verify.ScimUsers) + } + } +} + +func testScimGroupCreateValidation(t *testing.T, ds *Datastore) { + // Test validation for ExternalID + longString := strings.Repeat("a", SCIMMaxFieldLength+1) // String longer than allowed + + // Test ExternalID validation + groupWithLongExternalID := fleet.ScimGroup{ + DisplayName: "Valid Name", + ExternalID: ptr.String(longString), + ScimUsers: []uint{}, + } + _, err := ds.CreateScimGroup(context.Background(), &groupWithLongExternalID) + assert.Error(t, err) + assert.Contains(t, err.Error(), "external_id exceeds maximum length") + + // Test DisplayName validation + groupWithLongDisplayName := fleet.ScimGroup{ + DisplayName: longString, + ExternalID: ptr.String("valid-external-id"), + ScimUsers: []uint{}, + } + _, err = ds.CreateScimGroup(context.Background(), &groupWithLongDisplayName) + assert.Error(t, err) + assert.Contains(t, err.Error(), "display_name exceeds maximum length") + + // Test with valid values + validGroup := fleet.ScimGroup{ + DisplayName: "Valid Name", + ExternalID: ptr.String("valid-external-id"), + ScimUsers: []uint{}, + } + _, err = ds.CreateScimGroup(context.Background(), &validGroup) + assert.NoError(t, err) +} + +func testScimGroupByID(t *testing.T, ds *Datastore) { + // Create test users first + users := createTestScimUsers(t, ds) + userIDs := make([]uint, len(users)) + for i, user := range users { + userIDs[i] = user.ID + } + + // Create test groups + groups := createTestScimGroups(t, ds, userIDs) + + // Test retrieving each group + for _, tt := range groups { + returned, err := ds.ScimGroupByID(context.Background(), tt.ID) + assert.Nil(t, err) + assert.Equal(t, tt.ID, returned.ID) + assert.Equal(t, tt.DisplayName, returned.DisplayName) + assert.Equal(t, tt.ExternalID, returned.ExternalID) + + // Verify users + assert.Equal(t, len(tt.ScimUsers), len(returned.ScimUsers)) + if len(tt.ScimUsers) > 0 { + // Sort the user IDs for comparison + sort.Slice(tt.ScimUsers, func(i, j int) bool { + return tt.ScimUsers[i] < tt.ScimUsers[j] + }) + sort.Slice(returned.ScimUsers, func(i, j int) bool { + return returned.ScimUsers[i] < returned.ScimUsers[j] + }) + assert.Equal(t, tt.ScimUsers, returned.ScimUsers) + } + } + + // Test missing group + _, err := ds.ScimGroupByID(context.Background(), 10000000000) + assert.True(t, fleet.IsNotFound(err)) +} + +func testScimGroupByDisplayName(t *testing.T, ds *Datastore) { + // Create test users first + users := createTestScimUsers(t, ds) + userIDs := make([]uint, len(users)) + for i, user := range users { + userIDs[i] = user.ID + } + + // Create test groups + groups := createTestScimGroups(t, ds, userIDs) + + // Test retrieving each group by display name + for _, tt := range groups { + returned, err := ds.ScimGroupByDisplayName(context.Background(), tt.DisplayName) + assert.Nil(t, err) + assert.Equal(t, tt.ID, returned.ID) + assert.Equal(t, tt.DisplayName, returned.DisplayName) + assert.Equal(t, tt.ExternalID, returned.ExternalID) + + // Verify users + assert.Equal(t, len(tt.ScimUsers), len(returned.ScimUsers)) + if len(tt.ScimUsers) > 0 { + // Sort the user IDs for comparison + sort.Slice(tt.ScimUsers, func(i, j int) bool { + return tt.ScimUsers[i] < tt.ScimUsers[j] + }) + sort.Slice(returned.ScimUsers, func(i, j int) bool { + return returned.ScimUsers[i] < returned.ScimUsers[j] + }) + assert.Equal(t, tt.ScimUsers, returned.ScimUsers) + } + } + + // Test missing group + _, err := ds.ScimGroupByDisplayName(context.Background(), "Nonexistent Group") + assert.True(t, fleet.IsNotFound(err)) +} + +// createTestScimGroups creates test SCIM groups for testing +func createTestScimGroups(t *testing.T, ds *Datastore, userIDs []uint) []*fleet.ScimGroup { + createGroups := []fleet.ScimGroup{ + { + DisplayName: "Test Group 1", + ExternalID: ptr.String("ext-test-group-123"), + ScimUsers: []uint{}, + }, + { + DisplayName: "Test Group 2", + ExternalID: ptr.String("ext-test-group-456"), + ScimUsers: []uint{userIDs[0]}, + }, + { + DisplayName: "Test Group 3", + ExternalID: ptr.String("ext-test-group-789"), + ScimUsers: userIDs, + }, + } + + var groups []*fleet.ScimGroup + for _, g := range createGroups { + var err error + g.ID, err = ds.CreateScimGroup(context.Background(), &g) + require.Nil(t, err) + groups = append(groups, &g) + } + return groups +} + +func testReplaceScimGroup(t *testing.T, ds *Datastore) { + // Create test users first + users := createTestScimUsers(t, ds) + userIDs := make([]uint, len(users)) + for i, user := range users { + userIDs[i] = user.ID + } + + // Create a test group + group := fleet.ScimGroup{ + DisplayName: "Replace Test Group", + ExternalID: ptr.String("ext-replace-group-123"), + ScimUsers: []uint{userIDs[0]}, + } + + var err error + group.ID, err = ds.CreateScimGroup(context.Background(), &group) + require.Nil(t, err) + + // Verify the group was created correctly + createdGroup, err := ds.ScimGroupByID(context.Background(), group.ID) + require.Nil(t, err) + assert.Equal(t, group.DisplayName, createdGroup.DisplayName) + assert.Equal(t, group.ExternalID, createdGroup.ExternalID) + assert.Equal(t, 1, len(createdGroup.ScimUsers)) + assert.Equal(t, userIDs[0], createdGroup.ScimUsers[0]) + + // Modify the group + updatedGroup := fleet.ScimGroup{ + ID: group.ID, + DisplayName: "Updated Group", + ExternalID: ptr.String("ext-replace-group-456"), + ScimUsers: userIDs, // Add all users + } + + // Replace the group + err = ds.ReplaceScimGroup(context.Background(), &updatedGroup) + require.Nil(t, err) + + // Verify the group was updated correctly + replacedGroup, err := ds.ScimGroupByID(context.Background(), group.ID) + require.Nil(t, err) + assert.Equal(t, updatedGroup.DisplayName, replacedGroup.DisplayName) + assert.Equal(t, updatedGroup.ExternalID, replacedGroup.ExternalID) + + // Verify users were updated + assert.Equal(t, len(userIDs), len(replacedGroup.ScimUsers)) + + // Sort the user IDs for comparison + sort.Slice(userIDs, func(i, j int) bool { + return userIDs[i] < userIDs[j] + }) + sort.Slice(replacedGroup.ScimUsers, func(i, j int) bool { + return replacedGroup.ScimUsers[i] < replacedGroup.ScimUsers[j] + }) + assert.Equal(t, userIDs, replacedGroup.ScimUsers) + + // Test replacing a non-existent group + nonExistentGroup := fleet.ScimGroup{ + ID: 99999, // Non-existent ID + DisplayName: "Non-existent", + ExternalID: ptr.String("ext-non-existent"), + ScimUsers: []uint{}, + } + + err = ds.ReplaceScimGroup(context.Background(), &nonExistentGroup) + assert.True(t, fleet.IsNotFound(err)) +} + +func testScimGroupReplaceValidation(t *testing.T, ds *Datastore) { + // Create a valid group first + group := fleet.ScimGroup{ + DisplayName: "Validation Test Group", + ExternalID: ptr.String("ext-validation-group"), + ScimUsers: []uint{}, + } + + var err error + group.ID, err = ds.CreateScimGroup(context.Background(), &group) + require.NoError(t, err) + + // Test validation for ExternalID + longString := strings.Repeat("a", SCIMMaxFieldLength+1) // String longer than allowed + + // Test ExternalID validation + groupWithLongExternalID := fleet.ScimGroup{ + ID: group.ID, + DisplayName: "Valid Name", + ExternalID: ptr.String(longString), + ScimUsers: []uint{}, + } + err = ds.ReplaceScimGroup(context.Background(), &groupWithLongExternalID) + assert.Error(t, err) + assert.Contains(t, err.Error(), "external_id exceeds maximum length") + + // Test DisplayName validation + groupWithLongDisplayName := fleet.ScimGroup{ + ID: group.ID, + DisplayName: longString, + ExternalID: ptr.String("valid-external-id"), + ScimUsers: []uint{}, + } + err = ds.ReplaceScimGroup(context.Background(), &groupWithLongDisplayName) + assert.Error(t, err) + assert.Contains(t, err.Error(), "display_name exceeds maximum length") + + // Test with valid values + validGroup := fleet.ScimGroup{ + ID: group.ID, + DisplayName: "Updated Valid Name", + ExternalID: ptr.String("updated-valid-external-id"), + ScimUsers: []uint{}, + } + err = ds.ReplaceScimGroup(context.Background(), &validGroup) + assert.NoError(t, err) +} + +func testDeleteScimGroup(t *testing.T, ds *Datastore) { + // Create test users first + users := createTestScimUsers(t, ds) + userIDs := make([]uint, len(users)) + for i, user := range users { + userIDs[i] = user.ID + } + + // Create a test group + group := fleet.ScimGroup{ + DisplayName: "Delete Test Group", + ExternalID: ptr.String("ext-delete-group"), + ScimUsers: userIDs, + } + + var err error + group.ID, err = ds.CreateScimGroup(context.Background(), &group) + require.Nil(t, err) + + // Verify the group was created correctly + createdGroup, err := ds.ScimGroupByID(context.Background(), group.ID) + require.Nil(t, err) + assert.Equal(t, group.DisplayName, createdGroup.DisplayName) + + // Delete the group + err = ds.DeleteScimGroup(context.Background(), group.ID) + require.Nil(t, err) + + // Verify the group was deleted + _, err = ds.ScimGroupByID(context.Background(), group.ID) + assert.True(t, fleet.IsNotFound(err)) + + // Test deleting a non-existent group + err = ds.DeleteScimGroup(context.Background(), 99999) // Non-existent ID + assert.True(t, fleet.IsNotFound(err)) +} + +func testListScimGroups(t *testing.T, ds *Datastore) { + // Create test users first + users := createTestScimUsers(t, ds) + userIDs := make([]uint, len(users)) + for i, user := range users { + userIDs[i] = user.ID + } + + // Create test groups + groups := []fleet.ScimGroup{ + { + DisplayName: "List Test Group 1", + ExternalID: ptr.String("ext-list-group-123"), + ScimUsers: []uint{}, + }, + { + DisplayName: "List Test Group 2", + ExternalID: ptr.String("ext-list-group-456"), + ScimUsers: []uint{userIDs[0]}, + }, + { + DisplayName: "List Test Group 3", + ExternalID: ptr.String("ext-list-group-789"), + ScimUsers: userIDs, + }, + } + + // Create the groups + for i := range groups { + var err error + groups[i].ID, err = ds.CreateScimGroup(context.Background(), &groups[i]) + require.Nil(t, err) + } + + // Test 1: List all groups + allGroups, totalResults, err := ds.ListScimGroups(context.Background(), fleet.ScimListOptions{ + Page: 1, + PerPage: 10, + }) + require.Nil(t, err) + assert.GreaterOrEqual(t, len(allGroups), 3) // There might be other groups from previous tests + assert.GreaterOrEqual(t, totalResults, uint(3)) + + // Verify that our test groups are in the results + foundGroups := 0 + for _, g := range allGroups { + for _, testGroup := range groups { + if g.ID == testGroup.ID { + foundGroups++ + break + } + } + } + assert.Equal(t, 3, foundGroups) + + // Test 2: Pagination - first page with 2 items + page1Groups, totalPage1, err := ds.ListScimGroups(context.Background(), fleet.ScimListOptions{ + Page: 1, + PerPage: 2, + }) + require.Nil(t, err) + assert.Equal(t, 2, len(page1Groups)) + assert.GreaterOrEqual(t, totalPage1, uint(3)) // Total should be at least 3 + + // Test 3: Pagination - second page with 2 items + page2Groups, totalPage2, err := ds.ListScimGroups(context.Background(), fleet.ScimListOptions{ + Page: 2, + PerPage: 2, + }) + require.Nil(t, err) + assert.GreaterOrEqual(t, len(page2Groups), 1) // At least 1 item on the second page + assert.GreaterOrEqual(t, totalPage2, uint(3)) // Total should be at least 3 + + // Verify that page1 and page2 contain different groups + for _, p1Group := range page1Groups { + for _, p2Group := range page2Groups { + assert.NotEqual(t, p1Group.ID, p2Group.ID, "Groups should not appear on multiple pages") + } + } +} + +func testScimUserCreateValidation(t *testing.T, ds *Datastore) { + // Test validation for ExternalID + longString := strings.Repeat("a", SCIMMaxFieldLength+1) // String longer than SCIMMaxFieldLength + + // Test ExternalID validation + userWithLongExternalID := fleet.ScimUser{ + UserName: "valid-username", + ExternalID: ptr.String(longString), + GivenName: ptr.String("Valid"), + FamilyName: ptr.String("Name"), + Active: ptr.Bool(true), + } + _, err := ds.CreateScimUser(context.Background(), &userWithLongExternalID) + assert.Error(t, err) + assert.Contains(t, err.Error(), "external_id exceeds maximum length") + + // Test UserName validation + userWithLongUserName := fleet.ScimUser{ + UserName: longString, + ExternalID: ptr.String("valid-external-id"), + GivenName: ptr.String("Valid"), + FamilyName: ptr.String("Name"), + Active: ptr.Bool(true), + } + _, err = ds.CreateScimUser(context.Background(), &userWithLongUserName) + assert.Error(t, err) + assert.Contains(t, err.Error(), "user_name exceeds maximum length") + + // Test GivenName validation + userWithLongGivenName := fleet.ScimUser{ + UserName: "valid-username", + ExternalID: ptr.String("valid-external-id"), + GivenName: ptr.String(longString), + FamilyName: ptr.String("Name"), + Active: ptr.Bool(true), + } + _, err = ds.CreateScimUser(context.Background(), &userWithLongGivenName) + assert.Error(t, err) + assert.Contains(t, err.Error(), "given_name exceeds maximum length") + + // Test FamilyName validation + userWithLongFamilyName := fleet.ScimUser{ + UserName: "valid-username", + ExternalID: ptr.String("valid-external-id"), + GivenName: ptr.String("Valid"), + FamilyName: ptr.String(longString), + Active: ptr.Bool(true), + } + _, err = ds.CreateScimUser(context.Background(), &userWithLongFamilyName) + assert.Error(t, err) + assert.Contains(t, err.Error(), "family_name exceeds maximum length") + + // Test with valid values + validUser := fleet.ScimUser{ + UserName: "valid-username", + ExternalID: ptr.String("valid-external-id"), + GivenName: ptr.String("Valid"), + FamilyName: ptr.String("Name"), + Active: ptr.Bool(true), + } + _, err = ds.CreateScimUser(context.Background(), &validUser) + assert.NoError(t, err) +} + +func testScimUserReplaceValidation(t *testing.T, ds *Datastore) { + // Create a valid user first + user := fleet.ScimUser{ + UserName: "replace-validation-user", + ExternalID: ptr.String("ext-replace-validation"), + GivenName: ptr.String("Original"), + FamilyName: ptr.String("User"), + Active: ptr.Bool(true), + } + + var err error + user.ID, err = ds.CreateScimUser(context.Background(), &user) + require.NoError(t, err) + + // Test validation for ExternalID + longString := strings.Repeat("a", SCIMMaxFieldLength+1) // String longer than SCIMMaxFieldLength + + // Test ExternalID validation + userWithLongExternalID := fleet.ScimUser{ + ID: user.ID, + UserName: "valid-username", + ExternalID: ptr.String(longString), + GivenName: ptr.String("Valid"), + FamilyName: ptr.String("Name"), + Active: ptr.Bool(true), + } + err = ds.ReplaceScimUser(context.Background(), &userWithLongExternalID) + assert.Error(t, err) + assert.Contains(t, err.Error(), "external_id exceeds maximum length") + + // Test UserName validation + userWithLongUserName := fleet.ScimUser{ + ID: user.ID, + UserName: longString, + ExternalID: ptr.String("valid-external-id"), + GivenName: ptr.String("Valid"), + FamilyName: ptr.String("Name"), + Active: ptr.Bool(true), + } + err = ds.ReplaceScimUser(context.Background(), &userWithLongUserName) + assert.Error(t, err) + assert.Contains(t, err.Error(), "user_name exceeds maximum length") + + // Test GivenName validation + userWithLongGivenName := fleet.ScimUser{ + ID: user.ID, + UserName: "valid-username", + ExternalID: ptr.String("valid-external-id"), + GivenName: ptr.String(longString), + FamilyName: ptr.String("Name"), + Active: ptr.Bool(true), + } + err = ds.ReplaceScimUser(context.Background(), &userWithLongGivenName) + assert.Error(t, err) + assert.Contains(t, err.Error(), "given_name exceeds maximum length") + + // Test FamilyName validation + userWithLongFamilyName := fleet.ScimUser{ + ID: user.ID, + UserName: "valid-username", + ExternalID: ptr.String("valid-external-id"), + GivenName: ptr.String("Valid"), + FamilyName: ptr.String(longString), + Active: ptr.Bool(true), + } + err = ds.ReplaceScimUser(context.Background(), &userWithLongFamilyName) + assert.Error(t, err) + assert.Contains(t, err.Error(), "family_name exceeds maximum length") + + // Test with valid values + validUser := fleet.ScimUser{ + ID: user.ID, + UserName: "updated-username", + ExternalID: ptr.String("updated-external-id"), + GivenName: ptr.String("Updated"), + FamilyName: ptr.String("Name"), + Active: ptr.Bool(true), + } + err = ds.ReplaceScimUser(context.Background(), &validUser) + assert.NoError(t, err) +} diff --git a/server/fleet/datastore.go b/server/fleet/datastore.go index 4de568be07..a7a1949b30 100644 --- a/server/fleet/datastore.go +++ b/server/fleet/datastore.go @@ -2029,6 +2029,18 @@ type Datastore interface { DeleteScimUser(ctx context.Context, id uint) error // ListScimUsers retrieves a list of SCIM users with optional filtering ListScimUsers(ctx context.Context, opts ScimUsersListOptions) (users []ScimUser, totalResults uint, err error) + // CreateScimGroup creates a new SCIM group in the database + CreateScimGroup(ctx context.Context, group *ScimGroup) (uint, error) + // ScimGroupByID retrieves a SCIM group by ID + ScimGroupByID(ctx context.Context, id uint) (*ScimGroup, error) + // ScimGroupByDisplayName retrieves a SCIM group by display name + ScimGroupByDisplayName(ctx context.Context, displayName string) (*ScimGroup, error) + // ReplaceScimGroup replaces an existing SCIM group in the database + ReplaceScimGroup(ctx context.Context, group *ScimGroup) error + // DeleteScimGroup deletes a SCIM group from the database + DeleteScimGroup(ctx context.Context, id uint) error + // ListScimGroups retrieves a list of SCIM groups with pagination + ListScimGroups(ctx context.Context, opts ScimListOptions) (groups []ScimGroup, totalResults uint, err error) } type AndroidDatastore interface { diff --git a/server/fleet/scim.go b/server/fleet/scim.go index d9c6497553..f7e710f8bd 100644 --- a/server/fleet/scim.go +++ b/server/fleet/scim.go @@ -9,6 +9,7 @@ type ScimUser struct { FamilyName *string `db:"family_name"` Active *bool `db:"active"` Emails []ScimUserEmail + Groups []uint } func (su *ScimUser) AuthzType() string { @@ -23,11 +24,15 @@ type ScimUserEmail struct { Type *string `db:"type"` } -type ScimUsersListOptions struct { +type ScimListOptions struct { // Which page to return (must be positive integer) Page uint // How many results per page (must be positive integer) PerPage uint +} + +type ScimUsersListOptions struct { + ScimListOptions // UserNameFilter filters by userName -- max of 1 response is expected // Cannot be used with other filters. @@ -39,3 +44,10 @@ type ScimUsersListOptions struct { EmailTypeFilter *string EmailValueFilter *string } + +type ScimGroup struct { + ID uint `db:"id"` + ExternalID *string `db:"external_id"` + DisplayName string `db:"display_name"` + ScimUsers []uint +} diff --git a/server/mock/datastore_mock.go b/server/mock/datastore_mock.go index 038d407666..0821f709a9 100644 --- a/server/mock/datastore_mock.go +++ b/server/mock/datastore_mock.go @@ -1298,6 +1298,18 @@ type DeleteScimUserFunc func(ctx context.Context, id uint) error type ListScimUsersFunc func(ctx context.Context, opts fleet.ScimUsersListOptions) (users []fleet.ScimUser, totalResults uint, err error) +type CreateScimGroupFunc func(ctx context.Context, group *fleet.ScimGroup) (uint, error) + +type ScimGroupByIDFunc func(ctx context.Context, id uint) (*fleet.ScimGroup, error) + +type ScimGroupByDisplayNameFunc func(ctx context.Context, displayName string) (*fleet.ScimGroup, error) + +type ReplaceScimGroupFunc func(ctx context.Context, group *fleet.ScimGroup) error + +type DeleteScimGroupFunc func(ctx context.Context, id uint) error + +type ListScimGroupsFunc func(ctx context.Context, opts fleet.ScimListOptions) (groups []fleet.ScimGroup, totalResults uint, err error) + type DataStore struct { HealthCheckFunc HealthCheckFunc HealthCheckFuncInvoked bool @@ -3213,6 +3225,24 @@ type DataStore struct { ListScimUsersFunc ListScimUsersFunc ListScimUsersFuncInvoked bool + CreateScimGroupFunc CreateScimGroupFunc + CreateScimGroupFuncInvoked bool + + ScimGroupByIDFunc ScimGroupByIDFunc + ScimGroupByIDFuncInvoked bool + + ScimGroupByDisplayNameFunc ScimGroupByDisplayNameFunc + ScimGroupByDisplayNameFuncInvoked bool + + ReplaceScimGroupFunc ReplaceScimGroupFunc + ReplaceScimGroupFuncInvoked bool + + DeleteScimGroupFunc DeleteScimGroupFunc + DeleteScimGroupFuncInvoked bool + + ListScimGroupsFunc ListScimGroupsFunc + ListScimGroupsFuncInvoked bool + mu sync.Mutex } @@ -7681,3 +7711,45 @@ func (s *DataStore) ListScimUsers(ctx context.Context, opts fleet.ScimUsersListO s.mu.Unlock() return s.ListScimUsersFunc(ctx, opts) } + +func (s *DataStore) CreateScimGroup(ctx context.Context, group *fleet.ScimGroup) (uint, error) { + s.mu.Lock() + s.CreateScimGroupFuncInvoked = true + s.mu.Unlock() + return s.CreateScimGroupFunc(ctx, group) +} + +func (s *DataStore) ScimGroupByID(ctx context.Context, id uint) (*fleet.ScimGroup, error) { + s.mu.Lock() + s.ScimGroupByIDFuncInvoked = true + s.mu.Unlock() + return s.ScimGroupByIDFunc(ctx, id) +} + +func (s *DataStore) ScimGroupByDisplayName(ctx context.Context, displayName string) (*fleet.ScimGroup, error) { + s.mu.Lock() + s.ScimGroupByDisplayNameFuncInvoked = true + s.mu.Unlock() + return s.ScimGroupByDisplayNameFunc(ctx, displayName) +} + +func (s *DataStore) ReplaceScimGroup(ctx context.Context, group *fleet.ScimGroup) error { + s.mu.Lock() + s.ReplaceScimGroupFuncInvoked = true + s.mu.Unlock() + return s.ReplaceScimGroupFunc(ctx, group) +} + +func (s *DataStore) DeleteScimGroup(ctx context.Context, id uint) error { + s.mu.Lock() + s.DeleteScimGroupFuncInvoked = true + s.mu.Unlock() + return s.DeleteScimGroupFunc(ctx, id) +} + +func (s *DataStore) ListScimGroups(ctx context.Context, opts fleet.ScimListOptions) (groups []fleet.ScimGroup, totalResults uint, err error) { + s.mu.Lock() + s.ListScimGroupsFuncInvoked = true + s.mu.Unlock() + return s.ListScimGroupsFunc(ctx, opts) +}