Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@ require (
github.com/aws/aws-sdk-go-v2/service/secretsmanager v1.37.0
github.com/aws/aws-secretsmanager-caching-go/v2 v2.1.1
github.com/golang/mock v1.6.0
github.com/google/uuid v1.6.0
github.com/hashicorp/go-retryablehttp v0.7.8
github.com/pkg/errors v0.9.1
github.com/sirupsen/logrus v1.9.3
Expand Down Expand Up @@ -45,7 +46,6 @@ require (
github.com/go-logr/stdr v1.2.2 // indirect
github.com/go-viper/mapstructure/v2 v2.4.0 // indirect
github.com/google/s2a-go v0.1.9 // indirect
github.com/google/uuid v1.6.0 // indirect
github.com/googleapis/enterprise-certificate-proxy v0.3.6 // indirect
github.com/googleapis/gax-go/v2 v2.15.0 // indirect
github.com/hashicorp/go-cleanhttp v0.5.2 // indirect
Expand Down
4 changes: 4 additions & 0 deletions internal/aws/client_dry.go
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ func NewDryClient(c internal_http.Client, config *Config) (Client, error) {

func (dc *dryClient) CreateUser(u *interfaces.User) (*interfaces.User, error) {
log.WithField("user", u.Username).Info("DRY RUN: Would create user")
u.ID = virtualUserID(u.Username)
dc.virtualUsers[u.Username] = *u
return u, nil
}
Expand Down Expand Up @@ -61,6 +62,9 @@ func (dc *dryClient) FindUserByEmail(email string) (*interfaces.User, error) {

func (dc *dryClient) UpdateUser(u *interfaces.User) (*interfaces.User, error) {
log.WithField("user", u.Username).Info("DRY RUN: Would update user")
if u.ID == "" {
u.ID = virtualUserID(u.Username)
}
dc.virtualUsers[u.Username] = *u
return u, nil
}
108 changes: 108 additions & 0 deletions internal/aws/client_dry_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
package aws

import (
"errors"
"testing"

"github.com/awslabs/ssosync/internal/interfaces"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

// stubSCIMClient is a minimal stub of the SCIM Client interface for dry-client tests.
type stubSCIMClient struct {
findUserResult *interfaces.User
findUserErr error
}

func (s *stubSCIMClient) CreateUser(u *interfaces.User) (*interfaces.User, error) {
return u, nil
}
func (s *stubSCIMClient) FindGroupByDisplayName(name string) (*interfaces.Group, error) {
return nil, ErrGroupNotFound
}
func (s *stubSCIMClient) FindUserByEmail(email string) (*interfaces.User, error) {
return s.findUserResult, s.findUserErr
}
func (s *stubSCIMClient) UpdateUser(u *interfaces.User) (*interfaces.User, error) {
return u, nil
}

func newTestDryClient(t *testing.T, stub *stubSCIMClient) *dryClient {
t.Helper()
if stub == nil {
stub = &stubSCIMClient{findUserErr: ErrUserNotFound}
}
return &dryClient{
c: stub,
virtualUsers: make(map[string]interfaces.User),
}
}

func TestDryClient_CreateUser_PopulatesVirtualID(t *testing.T) {
dc := newTestDryClient(t, nil)
u := NewUser("Alice", "Smith", "alice@example.com", true)

result, err := dc.CreateUser(u)

require.NoError(t, err)
assert.Equal(t, virtualUserID("alice@example.com"), result.ID)
assert.True(t, isVirtualID(result.ID))
assert.Regexp(t, awsIDRegex, result.ID)
stored := dc.virtualUsers["alice@example.com"]
assert.Equal(t, result.ID, stored.ID)
}

func TestDryClient_CreateUser_Deterministic(t *testing.T) {
dc := newTestDryClient(t, nil)
u1 := NewUser("Alice", "Smith", "alice@example.com", true)
u2 := NewUser("Alice", "Smith", "alice@example.com", true)

r1, _ := dc.CreateUser(u1)
r2, _ := dc.CreateUser(u2)

assert.Equal(t, r1.ID, r2.ID, "same email must produce the same virtual ID")
}

func TestDryClient_FindUserByEmail_ReturnsVirtualUserWithID(t *testing.T) {
dc := newTestDryClient(t, &stubSCIMClient{findUserErr: ErrUserNotFound})
u := NewUser("Bob", "Jones", "bob@example.com", true)
_, err := dc.CreateUser(u)
require.NoError(t, err)

found, err := dc.FindUserByEmail("bob@example.com")

require.NoError(t, err)
require.NotNil(t, found)
assert.Equal(t, virtualUserID("bob@example.com"), found.ID)
assert.True(t, isVirtualID(found.ID))
}

func TestDryClient_UpdateUser_PreservesExistingID(t *testing.T) {
dc := newTestDryClient(t, nil)
u := UpdateUser("real-aws-id-1234", "Carol", "White", "carol@example.com", false)

result, err := dc.UpdateUser(u)

require.NoError(t, err)
assert.Equal(t, "real-aws-id-1234", result.ID, "real ID must not be replaced")
}

func TestDryClient_UpdateUser_PopulatesVirtualIDWhenEmpty(t *testing.T) {
dc := newTestDryClient(t, nil)
u := NewUser("Dan", "Brown", "dan@example.com", true)

result, err := dc.UpdateUser(u)

require.NoError(t, err)
assert.Equal(t, virtualUserID("dan@example.com"), result.ID)
}

func TestDryClient_FindUserByEmail_ForwardsRealError(t *testing.T) {
dc := newTestDryClient(t, &stubSCIMClient{findUserErr: errors.New("network error")})

result, err := dc.FindUserByEmail("err@example.com")

assert.Nil(t, result)
assert.EqualError(t, err, "network error")
}
50 changes: 45 additions & 5 deletions internal/aws/identitystore_dry.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,11 @@ package aws

import (
"context"
"slices"

"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/service/identitystore"
"github.com/aws/aws-sdk-go-v2/service/identitystore/types"
"github.com/awslabs/ssosync/internal/interfaces"
log "github.com/sirupsen/logrus"
)
Expand All @@ -26,18 +28,19 @@ func NewDryIdentityStore(client interfaces.IdentityStoreAPI) interfaces.Identity
func (d *DryIdentityStore) CreateGroup(ctx context.Context, params *identitystore.CreateGroupInput, optFns ...func(*identitystore.Options)) (*identitystore.CreateGroupOutput, error) {
log.WithField("displayName", *params.DisplayName).Info("DRY RUN: Would create group")
return &identitystore.CreateGroupOutput{
GroupId: aws.String(*params.DisplayName + "-virtual"),
GroupId: aws.String(virtualGroupID(*params.DisplayName)),
IdentityStoreId: params.IdentityStoreId,
}, nil
}

func (d *DryIdentityStore) CreateGroupMembership(ctx context.Context, params *identitystore.CreateGroupMembershipInput, optFns ...func(*identitystore.Options)) (*identitystore.CreateGroupMembershipOutput, error) {
memberValue := memberIDValue(params.MemberId)
log.WithFields(log.Fields{
"groupId": *params.GroupId,
"userId": params.MemberId,
"userId": memberValue,
}).Info("DRY RUN: Would create group membership")
return &identitystore.CreateGroupMembershipOutput{
MembershipId: aws.String("virtual-membership-id"),
MembershipId: aws.String(virtualMembershipID(*params.GroupId, memberValue)),
IdentityStoreId: params.IdentityStoreId,
}, nil
}
Expand All @@ -57,15 +60,42 @@ func (d *DryIdentityStore) DeleteUser(ctx context.Context, params *identitystore
return &identitystore.DeleteUserOutput{}, nil
}

// GetGroupMembershipId short-circuits when either the group or member is virtual
// so the real client is never called with a synthetic ID.
func (d *DryIdentityStore) GetGroupMembershipId(ctx context.Context, params *identitystore.GetGroupMembershipIdInput, optFns ...func(*identitystore.Options)) (*identitystore.GetGroupMembershipIdOutput, error) {
memberValue := memberIDValue(params.MemberId)
if isVirtualID(*params.GroupId) || isVirtualID(memberValue) {
return &identitystore.GetGroupMembershipIdOutput{
MembershipId: aws.String(virtualMembershipID(*params.GroupId, memberValue)),
IdentityStoreId: params.IdentityStoreId,
}, nil
}
return d.client.GetGroupMembershipId(ctx, params, optFns...)
}

// IsMemberInGroups short-circuits when the member or any group is virtual —
// a virtual user was never actually added, so membership is always false.
func (d *DryIdentityStore) IsMemberInGroups(ctx context.Context, params *identitystore.IsMemberInGroupsInput, optFns ...func(*identitystore.Options)) (*identitystore.IsMemberInGroupsOutput, error) {
return d.client.IsMemberInGroups(ctx, params, optFns...)
memberValue := memberIDValue(params.MemberId)
if !isVirtualID(memberValue) && !slices.ContainsFunc(params.GroupIds, isVirtualID) {
return d.client.IsMemberInGroups(ctx, params, optFns...)
}
results := make([]types.GroupMembershipExistenceResult, len(params.GroupIds))
for i, gid := range params.GroupIds {
results[i] = types.GroupMembershipExistenceResult{
GroupId: aws.String(gid),
MembershipExists: false,
}
}
return &identitystore.IsMemberInGroupsOutput{Results: results}, nil
}

// ListGroupMemberships short-circuits when the group is virtual — it has no
// real memberships to enumerate.
func (d *DryIdentityStore) ListGroupMemberships(ctx context.Context, params *identitystore.ListGroupMembershipsInput, optFns ...func(*identitystore.Options)) (*identitystore.ListGroupMembershipsOutput, error) {
if isVirtualID(*params.GroupId) {
return &identitystore.ListGroupMembershipsOutput{}, nil
}
return d.client.ListGroupMemberships(ctx, params, optFns...)
}

Expand All @@ -80,7 +110,17 @@ func (d *DryIdentityStore) ListUsers(ctx context.Context, params *identitystore.
func (d *DryIdentityStore) CreateUser(ctx context.Context, params *identitystore.CreateUserInput, optFns ...func(*identitystore.Options)) (*identitystore.CreateUserOutput, error) {
log.WithField("userName", *params.UserName).Info("DRY RUN: Would create user")
return &identitystore.CreateUserOutput{
UserId: aws.String(*params.UserName + "-virtual"),
UserId: aws.String(virtualUserID(*params.UserName)),
IdentityStoreId: params.IdentityStoreId,
}, nil
}

func memberIDValue(mid types.MemberId) string {
if mid == nil {
return ""
}
if m, ok := mid.(*types.MemberIdMemberUserId); ok {
return m.Value
}
return ""
}
115 changes: 111 additions & 4 deletions internal/aws/identitystore_dry_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,12 +32,13 @@ func TestDryIdentityStore_CreateGroup(t *testing.T) {
DisplayName: aws.String("Test Group"),
}

// Should not call the underlying client
// Underlying client must not be called.
result, err := dryStore.CreateGroup(ctx, input)

require.NoError(t, err)
assert.NotNil(t, result)
assert.Equal(t, "Test Group-virtual", *result.GroupId)
assert.Equal(t, virtualGroupID("Test Group"), *result.GroupId)
assert.True(t, isVirtualID(*result.GroupId))
assert.Equal(t, input.IdentityStoreId, result.IdentityStoreId)
}

Expand All @@ -56,7 +57,8 @@ func TestDryIdentityStore_CreateGroupMembership(t *testing.T) {

require.NoError(t, err)
assert.NotNil(t, result)
assert.Equal(t, "virtual-membership-id", *result.MembershipId)
assert.Equal(t, virtualMembershipID("group-123", "user-123"), *result.MembershipId)
assert.True(t, isVirtualID(*result.MembershipId))
assert.Equal(t, input.IdentityStoreId, result.IdentityStoreId)
}

Expand Down Expand Up @@ -92,13 +94,118 @@ func TestDryIdentityStore_DeleteUser(t *testing.T) {
assert.NotNil(t, result)
}

func TestDryIdentityStore_IsMemberInGroups_VirtualUser_ShortCircuits(t *testing.T) {
mockClient := mocks.NewMockIdentityStoreAPI(t)
dryStore := NewDryIdentityStore(mockClient)

ctx := context.Background()
virtualUID := virtualUserID("newuser@example.com")
realGroupID := "12345678-1234-1234-1234-123456789012"

input := &identitystore.IsMemberInGroupsInput{
IdentityStoreId: aws.String("d-123456789"),
GroupIds: []string{realGroupID},
MemberId: &types.MemberIdMemberUserId{Value: virtualUID},
}

// The underlying client must NOT be called.
result, err := dryStore.IsMemberInGroups(ctx, input)

require.NoError(t, err)
require.Len(t, result.Results, 1)
assert.False(t, result.Results[0].MembershipExists, "virtual user must not be a member")
assert.Equal(t, realGroupID, *result.Results[0].GroupId)
}

func TestDryIdentityStore_IsMemberInGroups_VirtualGroup_ShortCircuits(t *testing.T) {
mockClient := mocks.NewMockIdentityStoreAPI(t)
dryStore := NewDryIdentityStore(mockClient)

ctx := context.Background()
virtualGID := virtualGroupID("NewGroup")
realUID := "12345678-1234-1234-1234-123456789012"

input := &identitystore.IsMemberInGroupsInput{
IdentityStoreId: aws.String("d-123456789"),
GroupIds: []string{virtualGID},
MemberId: &types.MemberIdMemberUserId{Value: realUID},
}

result, err := dryStore.IsMemberInGroups(ctx, input)

require.NoError(t, err)
require.Len(t, result.Results, 1)
assert.False(t, result.Results[0].MembershipExists)
}

func TestDryIdentityStore_IsMemberInGroups_RealIDs_PassesThrough(t *testing.T) {
mockClient := mocks.NewMockIdentityStoreAPI(t)
dryStore := NewDryIdentityStore(mockClient)

ctx := context.Background()
input := &identitystore.IsMemberInGroupsInput{
IdentityStoreId: aws.String("d-123456789"),
GroupIds: []string{"group-real"},
MemberId: &types.MemberIdMemberUserId{Value: "user-real"},
}
expected := &identitystore.IsMemberInGroupsOutput{
Results: []types.GroupMembershipExistenceResult{
{GroupId: aws.String("group-real"), MembershipExists: true},
},
}

mockClient.EXPECT().IsMemberInGroups(ctx, input).Return(expected, nil).Once()

result, err := dryStore.IsMemberInGroups(ctx, input)
require.NoError(t, err)
assert.Equal(t, expected, result)
}

func TestDryIdentityStore_GetGroupMembershipId_VirtualMember_ShortCircuits(t *testing.T) {
mockClient := mocks.NewMockIdentityStoreAPI(t)
dryStore := NewDryIdentityStore(mockClient)

ctx := context.Background()
virtualUID := virtualUserID("new@example.com")
groupID := "12345678-1234-1234-1234-123456789012"

input := &identitystore.GetGroupMembershipIdInput{
IdentityStoreId: aws.String("d-123456789"),
GroupId: aws.String(groupID),
MemberId: &types.MemberIdMemberUserId{Value: virtualUID},
}

// Underlying client must NOT be called.
result, err := dryStore.GetGroupMembershipId(ctx, input)

require.NoError(t, err)
assert.Equal(t, virtualMembershipID(groupID, virtualUID), *result.MembershipId)
}

func TestDryIdentityStore_ListGroupMemberships_VirtualGroup_ShortCircuits(t *testing.T) {
mockClient := mocks.NewMockIdentityStoreAPI(t)
dryStore := NewDryIdentityStore(mockClient)

ctx := context.Background()
input := &identitystore.ListGroupMembershipsInput{
IdentityStoreId: aws.String("d-123456789"),
GroupId: aws.String(virtualGroupID("NewGroup")),
}

// Underlying client must NOT be called.
result, err := dryStore.ListGroupMemberships(ctx, input)

require.NoError(t, err)
assert.Empty(t, result.GroupMemberships)
}

func TestDryIdentityStore_PassThroughMethods(t *testing.T) {
mockClient := mocks.NewMockIdentityStoreAPI(t)
dryStore := NewDryIdentityStore(mockClient)

ctx := context.Background()

// Test IsMemberInGroups passes through to underlying client
// Test IsMemberInGroups passes through for real IDs
isMemberInput := &identitystore.IsMemberInGroupsInput{
IdentityStoreId: aws.String("d-123456789"),
GroupIds: []string{"group-123"},
Expand Down
Loading