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
This commit is contained in:
Victor Lyuboslavsky
2025-04-02 17:10:40 -05:00
committed by GitHub
parent 28c964b687
commit 8658608c37
8 changed files with 1757 additions and 37 deletions
+327
View File
@@ -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
}
// Microsofts SCIM implementation (Entra ID) imposes additional constraints—like enforcing uniqueness on a groups
// displayName—that the SCIM spec itself does not mandate.
// In effect, Microsofts implementation diverges from strict SCIM compliance by making displayName behave like a unique key.
// SCIM only mandates that each groups "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-<id>'", 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
}
+68 -3
View File
@@ -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
+43 -12
View File
@@ -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
}
+474 -2
View File
@@ -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
}
+748 -19
View File
@@ -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)
}
+12
View File
@@ -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 {
+13 -1
View File
@@ -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
}
+72
View File
@@ -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)
}