Files

838 lines
26 KiB
Go

package scimsync
import (
"bytes"
"encoding/json"
"fmt"
"io"
"maps"
"net/http"
"net/url"
"slices"
"sort"
"strconv"
"strings"
"sync"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"github.com/pocket-id/pocket-id/backend/internal/model"
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
)
const (
mockSCIMEndpoint = "https://scim.example.test"
scimListResponseSchema = "urn:ietf:params:scim:api:messages:2.0:ListResponse"
scimErrorResponseSchema = "urn:ietf:params:scim:api:messages:2.0:Error"
mockSCIMRequestContentType = "application/scim+json"
)
type scimSyncFixture struct {
db *gorm.DB
service *Service
transport *mockSCIMTransport
client model.OidcClient
provider ServiceProvider
}
func newSCIMSyncFixture(t *testing.T, restricted bool) *scimSyncFixture {
t.Helper()
db := testutils.NewDatabaseForTest(t)
providerToken := t.Name()
client := model.OidcClient{
Base: model.Base{ID: "oidc-client"},
Name: "SCIM client",
IsGroupRestricted: restricted,
}
err := db.Create(&client).Error
require.NoError(t, err)
provider := ServiceProvider{
Base: model.Base{ID: "scim-provider"},
Endpoint: mockSCIMEndpoint,
Token: datatype.EncryptedString(providerToken),
OidcClientID: client.ID,
}
err = db.Create(&provider).Error
require.NoError(t, err)
transport := newMockSCIMTransport(providerToken)
service := newService(db, &http.Client{Transport: transport})
return &scimSyncFixture{
db: db,
service: service,
transport: transport,
client: client,
provider: provider,
}
}
func (f *scimSyncFixture) createUser(t *testing.T, id, username string, email *string, disabled bool) model.User {
t.Helper()
user := model.User{
Base: model.Base{ID: id},
Username: username,
Email: email,
FirstName: strings.ToUpper(username[:1]) + username[1:],
LastName: "Example",
DisplayName: username + " display",
Disabled: disabled,
}
err := f.db.Create(&user).Error
require.NoError(t, err)
return user
}
func (f *scimSyncFixture) createGroup(t *testing.T, id, name string, users ...model.User) model.UserGroup {
t.Helper()
group := model.UserGroup{
Base: model.Base{ID: id},
Name: name,
FriendlyName: name + " friendly",
}
err := f.db.Create(&group).Error
require.NoError(t, err)
if len(users) > 0 {
err = f.db.Model(&group).Association("Users").Replace(users)
require.NoError(t, err)
}
return group
}
func (f *scimSyncFixture) allowGroups(t *testing.T, groups ...model.UserGroup) {
t.Helper()
err := f.db.Model(&f.client).Association("AllowedUserGroups").Replace(groups)
require.NoError(t, err)
}
func (f *scimSyncFixture) requireLastSynced(t *testing.T, expected bool) {
t.Helper()
var provider ServiceProvider
err := f.db.First(&provider, "id = ?", f.provider.ID).Error
require.NoError(t, err)
if expected {
require.NotNil(t, provider.LastSyncedAt)
} else {
require.Nil(t, provider.LastSyncedAt)
}
}
func TestSyncCreatesCompliantUsersBeforeGroups(t *testing.T) {
fixture := newSCIMSyncFixture(t, false)
aliceEmail := "alice@example.com"
alice := fixture.createUser(t, "user-alice", "alice", &aliceEmail, false)
bob := fixture.createUser(t, "user-bob", "bob", nil, true)
group := fixture.createGroup(t, "group-engineering", "engineering", alice, bob)
require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID))
users := fixture.transport.usersSnapshot()
require.Len(t, users, 2)
remoteAlice := resourceByExternalID(alice.ID, users)
require.NotNil(t, remoteAlice)
assert.Equal(t, "alice", remoteAlice.UserName)
assert.Equal(t, "Alice", remoteAlice.Name.GivenName)
assert.Equal(t, "Example", remoteAlice.Name.FamilyName)
assert.Equal(t, alice.DisplayName, remoteAlice.Display)
assert.True(t, remoteAlice.Active)
require.Equal(t, []ScimEmail{{Value: aliceEmail, Primary: true}}, remoteAlice.Emails)
remoteBob := resourceByExternalID(bob.ID, users)
require.NotNil(t, remoteBob)
assert.False(t, remoteBob.Active)
assert.Empty(t, remoteBob.Emails)
groups := fixture.transport.groupsSnapshot()
require.Len(t, groups, 1)
remoteGroup := resourceByExternalID(group.ID, groups)
require.NotNil(t, remoteGroup)
assert.Equal(t, group.FriendlyName, remoteGroup.Display)
assert.ElementsMatch(t, []ScimGroupMember{{Value: remoteAlice.ID}, {Value: remoteBob.ID}}, remoteGroup.Members)
requests := fixture.transport.requestsSnapshot()
lastUserCreate := lastRequestIndex(requests, http.MethodPost, "/Users")
firstGroupCreate := firstRequestIndex(requests, http.MethodPost, "/Groups")
require.NotEqual(t, -1, lastUserCreate)
require.Greater(t, firstGroupCreate, lastUserCreate)
fixture.requireLastSynced(t, true)
fixture.transport.requireCompliant(t)
}
func TestSyncUpdatesExistingUsersAndGroupsWithPUT(t *testing.T) {
fixture := newSCIMSyncFixture(t, false)
email := "updated@example.com"
user := fixture.createUser(t, "user-updated", "updated", &email, true)
group := fixture.createGroup(t, "group-updated", "updated", user)
remoteModified := time.Now().Add(-time.Hour)
fixture.transport.seedUser(ScimUser{
ScimResourceData: remoteResourceData("remote-user", user.ID, scimUserSchema, "User", remoteModified),
UserName: "stale-name",
Active: true,
})
fixture.transport.seedGroup(ScimGroup{
ScimResourceData: remoteResourceData("remote-group", group.ID, scimGroupSchema, "Group", remoteModified),
Display: "stale-group",
})
require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID))
remoteUser := fixture.transport.user("remote-user")
require.NotNil(t, remoteUser)
assert.Equal(t, user.Username, remoteUser.UserName)
assert.Equal(t, user.DisplayName, remoteUser.Display)
assert.False(t, remoteUser.Active)
require.Equal(t, []ScimEmail{{Value: email, Primary: true}}, remoteUser.Emails)
remoteGroup := fixture.transport.group("remote-group")
require.NotNil(t, remoteGroup)
assert.Equal(t, group.FriendlyName, remoteGroup.Display)
require.Equal(t, []ScimGroupMember{{Value: remoteUser.ID}}, remoteGroup.Members)
requests := fixture.transport.requestsSnapshot()
assert.Equal(t, 1, countRequests(requests, http.MethodPut, "/Users/remote-user"))
assert.Equal(t, 1, countRequests(requests, http.MethodPut, "/Groups/remote-group"))
assert.Zero(t, countRequestsWithPrefix(requests, http.MethodPost, "/"))
fixture.requireLastSynced(t, true)
fixture.transport.requireCompliant(t)
}
func TestSyncRestrictedClientDeletesDisallowedResources(t *testing.T) {
fixture := newSCIMSyncFixture(t, true)
allowedUser := fixture.createUser(t, "user-allowed", "allowed", nil, false)
deniedUser := fixture.createUser(t, "user-denied", "denied", nil, false)
allowedGroup := fixture.createGroup(t, "group-allowed", "allowed", allowedUser)
deniedGroup := fixture.createGroup(t, "group-denied", "denied", deniedUser)
fixture.allowGroups(t, allowedGroup)
remoteModified := time.Now().Add(time.Hour)
fixture.transport.seedUser(ScimUser{
ScimResourceData: remoteResourceData("remote-allowed-user", allowedUser.ID, scimUserSchema, "User", remoteModified),
UserName: allowedUser.Username,
Active: true,
})
fixture.transport.seedUser(ScimUser{
ScimResourceData: remoteResourceData("remote-denied-user", deniedUser.ID, scimUserSchema, "User", remoteModified),
UserName: deniedUser.Username,
Active: true,
})
fixture.transport.seedGroup(ScimGroup{
ScimResourceData: remoteResourceData("remote-allowed-group", allowedGroup.ID, scimGroupSchema, "Group", remoteModified),
Display: allowedGroup.FriendlyName,
Members: []ScimGroupMember{{Value: "remote-allowed-user"}},
})
fixture.transport.seedGroup(ScimGroup{
ScimResourceData: remoteResourceData("remote-denied-group", deniedGroup.ID, scimGroupSchema, "Group", remoteModified),
Display: deniedGroup.FriendlyName,
Members: []ScimGroupMember{{Value: "remote-denied-user"}},
})
require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID))
assert.NotNil(t, fixture.transport.user("remote-allowed-user"))
assert.Nil(t, fixture.transport.user("remote-denied-user"))
assert.NotNil(t, fixture.transport.group("remote-allowed-group"))
assert.Nil(t, fixture.transport.group("remote-denied-group"))
requests := fixture.transport.requestsSnapshot()
assert.Equal(t, 1, countRequests(requests, http.MethodDelete, "/Users/remote-denied-user"))
assert.Equal(t, 1, countRequests(requests, http.MethodDelete, "/Groups/remote-denied-group"))
fixture.requireLastSynced(t, true)
fixture.transport.requireCompliant(t)
}
func TestSyncSkipsResourcesNewerThanTheLocalSnapshotAndPaginates(t *testing.T) {
fixture := newSCIMSyncFixture(t, false)
alice := fixture.createUser(t, "user-alice", "alice", nil, false)
bob := fixture.createUser(t, "user-bob", "bob", nil, false)
remoteModified := time.Now().Add(time.Hour)
fixture.transport.pageSize = 1
fixture.transport.seedUser(ScimUser{
ScimResourceData: remoteResourceData("remote-alice", alice.ID, scimUserSchema, "User", remoteModified),
UserName: alice.Username,
Active: true,
})
fixture.transport.seedUser(ScimUser{
ScimResourceData: remoteResourceData("remote-bob", bob.ID, scimUserSchema, "User", remoteModified),
UserName: bob.Username,
Active: true,
})
require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID))
requests := fixture.transport.requestsSnapshot()
assert.Equal(t, []string{"1", "2"}, queryValues(requests, http.MethodGet, "/Users", "startIndex"))
assert.Equal(t, []string{"1000", "1000"}, queryValues(requests, http.MethodGet, "/Users", "count"))
assert.Zero(t, countMutationRequests(requests))
fixture.requireLastSynced(t, true)
fixture.transport.requireCompliant(t)
}
func TestSyncContinuesAfterResourceFailureAndDoesNotMarkCompletion(t *testing.T) {
fixture := newSCIMSyncFixture(t, false)
fixture.createUser(t, "user-success", "success", nil, false)
fixture.createUser(t, "user-failure", "failure", nil, false)
fixture.transport.failCreates["user-failure"] = http.StatusInternalServerError
err := fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID)
require.Error(t, err)
require.ErrorContains(t, err, "status 500")
users := fixture.transport.usersSnapshot()
assert.NotNil(t, resourceByExternalID("user-success", users))
assert.Nil(t, resourceByExternalID("user-failure", users))
fixture.requireLastSynced(t, false)
fixture.transport.requireCompliant(t)
}
func TestSyncRetriesRateLimitedSCIMRequests(t *testing.T) {
fixture := newSCIMSyncFixture(t, false)
fixture.transport.rateLimits[http.MethodGet+" /Users"] = 2
require.NoError(t, fixture.service.SyncServiceProvider(t.Context(), fixture.provider.ID))
requests := fixture.transport.requestsSnapshot()
assert.Equal(t, 3, countRequests(requests, http.MethodGet, "/Users"))
fixture.requireLastSynced(t, true)
fixture.transport.requireCompliant(t)
}
func remoteResourceData(id, externalID, schema, resourceType string, modified time.Time) ScimResourceData {
return ScimResourceData{
ID: id,
ExternalID: externalID,
Schemas: []string{schema},
Meta: &ScimResourceMeta{
Location: mockSCIMEndpoint + "/" + resourceType + "s/" + id,
ResourceType: resourceType,
Created: modified.Add(-time.Hour),
LastModified: modified,
Version: `W/"seed"`,
},
}
}
func resourceByExternalID[T ScimResource](externalID string, resources map[string]T) *T {
for _, resource := range resources {
if resource.GetExternalID() == externalID {
result := resource
return &result
}
}
return nil
}
type mockSCIMRequest struct {
method string
path string
query url.Values
body []byte
}
// mockSCIMTransport is a *http.Transport that mocks HTTP server responses to test for compliance with SCIM specs
type mockSCIMTransport struct {
mu sync.Mutex
expectedToken string
pageSize int
nextID int
users map[string]ScimUser
groups map[string]ScimGroup
requests []mockSCIMRequest
violations []string
rateLimits map[string]int
failCreates map[string]int
}
func newMockSCIMTransport(expectedToken string) *mockSCIMTransport {
return &mockSCIMTransport{
expectedToken: expectedToken,
users: map[string]ScimUser{},
groups: map[string]ScimGroup{},
rateLimits: map[string]int{},
failCreates: map[string]int{},
}
}
func (m *mockSCIMTransport) RoundTrip(req *http.Request) (*http.Response, error) {
var body []byte
if req.Body != nil {
var err error
body, err = io.ReadAll(req.Body)
if err != nil {
return nil, err
}
}
m.mu.Lock()
defer m.mu.Unlock()
m.requests = append(m.requests, mockSCIMRequest{
method: req.Method,
path: req.URL.Path,
query: req.URL.Query(),
body: slices.Clone(body),
})
if violation := m.validateRequest(req, body); violation != "" {
m.violations = append(m.violations, violation)
return mockSCIMErrorResponse(req, http.StatusBadRequest, violation), nil
}
requestKey := req.Method + " " + req.URL.Path
if m.rateLimits[requestKey] > 0 {
m.rateLimits[requestKey]--
response := mockSCIMErrorResponse(req, http.StatusTooManyRequests, "rate limited")
response.Header.Set("Retry-After", "0")
return response, nil
}
segments := strings.Split(strings.Trim(req.URL.Path, "/"), "/")
if len(segments) == 0 || segments[0] == "" {
return mockSCIMErrorResponse(req, http.StatusNotFound, "resource path is empty"), nil
}
switch segments[0] {
case "Users":
return m.handleUsers(req, segments, body), nil
case "Groups":
return m.handleGroups(req, segments, body), nil
default:
return mockSCIMErrorResponse(req, http.StatusNotFound, "resource type is unknown"), nil
}
}
func (m *mockSCIMTransport) validateRequest(req *http.Request, body []byte) string {
if req.URL.Scheme != "https" || req.URL.Host != "scim.example.test" {
return fmt.Sprintf("request used unexpected SCIM endpoint %s", req.URL.String())
}
if req.Header.Get("Accept") != mockSCIMRequestContentType {
return "request did not accept application/scim+json"
}
if req.Header.Get("Authorization") != "Bearer "+m.expectedToken {
return "request did not use the configured bearer token"
}
if len(body) > 0 && req.Header.Get("Content-Type") != mockSCIMRequestContentType {
return "request body did not use application/scim+json"
}
return ""
}
func (m *mockSCIMTransport) handleUsers(req *http.Request, segments []string, body []byte) *http.Response {
switch req.Method {
case http.MethodGet:
if len(segments) != 1 {
return mockSCIMErrorResponse(req, http.StatusMethodNotAllowed, "individual user reads are unsupported")
}
resources := make([]ScimUser, 0, len(m.users))
for _, user := range m.users {
resources = append(resources, user)
}
sort.Slice(resources, func(i, j int) bool { return resources[i].ID < resources[j].ID })
return mockSCIMListResponse(req, resources, m.pageSize)
case http.MethodPost:
if len(segments) != 1 {
return mockSCIMErrorResponse(req, http.StatusNotFound, "user collection path is invalid")
}
var user ScimUser
err := json.Unmarshal(body, &user)
if err != nil {
return mockSCIMErrorResponse(req, http.StatusBadRequest, "user payload is invalid JSON")
}
violation := validateUserPayload(body, user)
if violation != "" {
m.violations = append(m.violations, violation)
return mockSCIMErrorResponse(req, http.StatusBadRequest, violation)
}
status := m.failCreates[user.ExternalID]
if status != 0 {
return mockSCIMErrorResponse(req, status, "injected user creation failure")
}
m.nextID++
user.ID = fmt.Sprintf("remote-user-%d", m.nextID)
user.Meta = newMockMeta("User", "/Users/"+user.ID, m.nextID)
m.users[user.ID] = user
return mockSCIMJSONResponse(req, http.StatusCreated, user)
case http.MethodPut:
if len(segments) != 2 {
return mockSCIMErrorResponse(req, http.StatusNotFound, "user resource path is invalid")
}
_, ok := m.users[segments[1]]
if !ok {
return mockSCIMErrorResponse(req, http.StatusNotFound, "user does not exist")
}
var user ScimUser
err := json.Unmarshal(body, &user)
if err != nil {
return mockSCIMErrorResponse(req, http.StatusBadRequest, "user payload is invalid JSON")
}
violation := validateUserPayload(body, user)
if violation != "" {
m.violations = append(m.violations, violation)
return mockSCIMErrorResponse(req, http.StatusBadRequest, violation)
}
m.nextID++
user.ID = segments[1]
user.Meta = newMockMeta("User", "/Users/"+user.ID, m.nextID)
m.users[user.ID] = user
return mockSCIMJSONResponse(req, http.StatusOK, user)
case http.MethodDelete:
if len(segments) != 2 {
return mockSCIMErrorResponse(req, http.StatusNotFound, "user resource path is invalid")
}
_, ok := m.users[segments[1]]
if !ok {
return mockSCIMErrorResponse(req, http.StatusNotFound, "user does not exist")
}
delete(m.users, segments[1])
return mockSCIMNoContentResponse(req)
default:
return mockSCIMErrorResponse(req, http.StatusMethodNotAllowed, "user method is unsupported")
}
}
func (m *mockSCIMTransport) handleGroups(req *http.Request, segments []string, body []byte) *http.Response {
switch req.Method {
case http.MethodGet:
if len(segments) != 1 {
return mockSCIMErrorResponse(req, http.StatusMethodNotAllowed, "individual group reads are unsupported")
}
resources := make([]ScimGroup, 0, len(m.groups))
for _, group := range m.groups {
resources = append(resources, group)
}
sort.Slice(resources, func(i, j int) bool { return resources[i].ID < resources[j].ID })
return mockSCIMListResponse(req, resources, m.pageSize)
case http.MethodPost:
if len(segments) != 1 {
return mockSCIMErrorResponse(req, http.StatusNotFound, "group collection path is invalid")
}
var group ScimGroup
err := json.Unmarshal(body, &group)
if err != nil {
return mockSCIMErrorResponse(req, http.StatusBadRequest, "group payload is invalid JSON")
}
violation := m.validateGroupPayload(body, group)
if violation != "" {
m.violations = append(m.violations, violation)
return mockSCIMErrorResponse(req, http.StatusBadRequest, violation)
}
status := m.failCreates[group.ExternalID]
if status != 0 {
return mockSCIMErrorResponse(req, status, "injected group creation failure")
}
m.nextID++
group.ID = fmt.Sprintf("remote-group-%d", m.nextID)
group.Meta = newMockMeta("Group", "/Groups/"+group.ID, m.nextID)
m.groups[group.ID] = group
return mockSCIMJSONResponse(req, http.StatusCreated, group)
case http.MethodPut:
if len(segments) != 2 {
return mockSCIMErrorResponse(req, http.StatusNotFound, "group resource path is invalid")
}
_, ok := m.groups[segments[1]]
if !ok {
return mockSCIMErrorResponse(req, http.StatusNotFound, "group does not exist")
}
var group ScimGroup
err := json.Unmarshal(body, &group)
if err != nil {
return mockSCIMErrorResponse(req, http.StatusBadRequest, "group payload is invalid JSON")
}
violation := m.validateGroupPayload(body, group)
if violation != "" {
m.violations = append(m.violations, violation)
return mockSCIMErrorResponse(req, http.StatusBadRequest, violation)
}
m.nextID++
group.ID = segments[1]
group.Meta = newMockMeta("Group", "/Groups/"+group.ID, m.nextID)
m.groups[group.ID] = group
return mockSCIMJSONResponse(req, http.StatusOK, group)
case http.MethodDelete:
if len(segments) != 2 {
return mockSCIMErrorResponse(req, http.StatusNotFound, "group resource path is invalid")
}
_, ok := m.groups[segments[1]]
if !ok {
return mockSCIMErrorResponse(req, http.StatusNotFound, "group does not exist")
}
delete(m.groups, segments[1])
return mockSCIMNoContentResponse(req)
default:
return mockSCIMErrorResponse(req, http.StatusMethodNotAllowed, "group method is unsupported")
}
}
func validateUserPayload(body []byte, user ScimUser) string {
if violation := validateResourcePayload(body, user.ScimResourceData, scimUserSchema); violation != "" {
return violation
}
if user.UserName == "" {
return "SCIM user payload omitted userName"
}
return ""
}
func (m *mockSCIMTransport) validateGroupPayload(body []byte, group ScimGroup) string {
if violation := validateResourcePayload(body, group.ScimResourceData, scimGroupSchema); violation != "" {
return violation
}
if group.Display == "" {
return "SCIM group payload omitted displayName"
}
for _, member := range group.Members {
if _, ok := m.users[member.Value]; !ok {
return fmt.Sprintf("SCIM group referenced unknown user %q", member.Value)
}
}
return ""
}
func validateResourcePayload(body []byte, resource ScimResourceData, expectedSchema string) string {
if resource.ExternalID == "" {
return "SCIM resource payload omitted externalId"
}
if !slices.Contains(resource.Schemas, expectedSchema) {
return fmt.Sprintf("SCIM resource payload omitted schema %q", expectedSchema)
}
var raw map[string]json.RawMessage
err := json.Unmarshal(body, &raw)
if err != nil {
return "SCIM resource payload is invalid JSON"
}
_, ok := raw["id"]
if ok {
return "SCIM write payload included read-only id"
}
_, ok = raw["meta"]
if ok {
return "SCIM write payload included read-only meta"
}
return ""
}
func mockSCIMListResponse[T any](req *http.Request, resources []T, pageSize int) *http.Response {
startIndex, err := strconv.Atoi(req.URL.Query().Get("startIndex"))
if err != nil || startIndex < 1 {
return mockSCIMErrorResponse(req, http.StatusBadRequest, "startIndex must be a one-based integer")
}
count, err := strconv.Atoi(req.URL.Query().Get("count"))
if err != nil || count < 1 {
return mockSCIMErrorResponse(req, http.StatusBadRequest, "count must be a positive integer")
}
if pageSize > 0 && pageSize < count {
count = pageSize
}
start := min(startIndex-1, len(resources))
end := min(start+count, len(resources))
page := resources[start:end]
return mockSCIMJSONResponse(req, http.StatusOK, struct {
Schemas []string `json:"schemas"`
Resources []T `json:"Resources"`
TotalResults int `json:"totalResults"`
StartIndex int `json:"startIndex"`
ItemsPerPage int `json:"itemsPerPage"`
}{
Schemas: []string{scimListResponseSchema},
Resources: page,
TotalResults: len(resources),
StartIndex: startIndex,
ItemsPerPage: len(page),
})
}
func mockSCIMJSONResponse(req *http.Request, status int, payload any) *http.Response {
body, err := json.Marshal(payload)
if err != nil {
panic(err)
}
return &http.Response{
StatusCode: status,
Header: http.Header{"Content-Type": []string{mockSCIMRequestContentType}},
Body: io.NopCloser(bytes.NewReader(body)),
ContentLength: int64(len(body)),
Request: req,
}
}
func mockSCIMErrorResponse(req *http.Request, status int, detail string) *http.Response {
return mockSCIMJSONResponse(req, status, struct {
Schemas []string `json:"schemas"`
Status string `json:"status"`
Detail string `json:"detail"`
}{
Schemas: []string{scimErrorResponseSchema},
Status: strconv.Itoa(status),
Detail: detail,
})
}
func mockSCIMNoContentResponse(req *http.Request) *http.Response {
return &http.Response{
StatusCode: http.StatusNoContent,
Header: make(http.Header),
Body: http.NoBody,
Request: req,
}
}
func newMockMeta(resourceType, resourcePath string, version int) *ScimResourceMeta {
now := time.Now().UTC()
return &ScimResourceMeta{
Location: mockSCIMEndpoint + resourcePath,
ResourceType: resourceType,
Created: now,
LastModified: now,
Version: fmt.Sprintf(`W/"%d"`, version),
}
}
func (m *mockSCIMTransport) seedUser(user ScimUser) {
m.mu.Lock()
defer m.mu.Unlock()
m.users[user.ID] = user
}
func (m *mockSCIMTransport) seedGroup(group ScimGroup) {
m.mu.Lock()
defer m.mu.Unlock()
m.groups[group.ID] = group
}
func (m *mockSCIMTransport) user(id string) *ScimUser {
m.mu.Lock()
defer m.mu.Unlock()
user, ok := m.users[id]
if !ok {
return nil
}
return &user
}
func (m *mockSCIMTransport) group(id string) *ScimGroup {
m.mu.Lock()
defer m.mu.Unlock()
group, ok := m.groups[id]
if !ok {
return nil
}
return &group
}
func (m *mockSCIMTransport) usersSnapshot() map[string]ScimUser {
m.mu.Lock()
defer m.mu.Unlock()
return cloneMap(m.users)
}
func (m *mockSCIMTransport) groupsSnapshot() map[string]ScimGroup {
m.mu.Lock()
defer m.mu.Unlock()
return cloneMap(m.groups)
}
func (m *mockSCIMTransport) requestsSnapshot() []mockSCIMRequest {
m.mu.Lock()
defer m.mu.Unlock()
return slices.Clone(m.requests)
}
func (m *mockSCIMTransport) requireCompliant(t *testing.T) {
t.Helper()
m.mu.Lock()
defer m.mu.Unlock()
require.Empty(t, m.violations)
}
func cloneMap[K comparable, V any](input map[K]V) map[K]V {
result := make(map[K]V, len(input))
maps.Copy(result, input)
return result
}
func firstRequestIndex(requests []mockSCIMRequest, method, path string) int {
for i, request := range requests {
if request.method == method && request.path == path {
return i
}
}
return -1
}
func lastRequestIndex(requests []mockSCIMRequest, method, path string) int {
for i := len(requests) - 1; i >= 0; i-- {
if requests[i].method == method && requests[i].path == path {
return i
}
}
return -1
}
func countRequests(requests []mockSCIMRequest, method, path string) int {
count := 0
for _, request := range requests {
if request.method == method && request.path == path {
count++
}
}
return count
}
func countRequestsWithPrefix(requests []mockSCIMRequest, method, pathPrefix string) int {
count := 0
for _, request := range requests {
if request.method == method && strings.HasPrefix(request.path, pathPrefix) {
count++
}
}
return count
}
func countMutationRequests(requests []mockSCIMRequest) int {
count := 0
for _, request := range requests {
if request.method == http.MethodPost || request.method == http.MethodPut || request.method == http.MethodDelete {
count++
}
}
return count
}
func queryValues(requests []mockSCIMRequest, method, path, key string) []string {
values := make([]string, 0)
for _, request := range requests {
if request.method == method && request.path == path {
values = append(values, request.query.Get(key))
}
}
return values
}