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:
@@ -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-<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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user