diff --git a/internal/entraidclient/client.go b/internal/entraidclient/client.go new file mode 100644 index 0000000..b016853 --- /dev/null +++ b/internal/entraidclient/client.go @@ -0,0 +1,136 @@ +package entraidclient + +import ( + "context" + "fmt" + + "cloud.google.com/go/auth/credentials/idtoken" + "github.com/Azure/azure-sdk-for-go/sdk/azidentity" + "github.com/google/uuid" + msgraphsdkgo "github.com/microsoftgraph/msgraph-sdk-go" + msgraphcore "github.com/microsoftgraph/msgraph-sdk-go-core" + "github.com/microsoftgraph/msgraph-sdk-go/models" +) + +type Client struct { + client *msgraphsdkgo.GraphServiceClient +} + +func New(tenantId, clientId string) (*Client, error) { + creds, err := azidentity.NewClientAssertionCredential(tenantId, clientId, func(ctx context.Context) (string, error) { + creds, err := idtoken.NewCredentials(&idtoken.Options{Audience: "api://AzureADTokenExchange"}) + if err != nil { + return "", err + } + token, err := creds.Token(ctx) + if err != nil { + return "", err + } + return token.Value, nil + }, nil) + if err != nil { + return nil, fmt.Errorf("exchange for azure credentials: %w", err) + } + + client, err := msgraphsdkgo.NewGraphServiceClientWithCredentials(creds, []string{"https://graph.microsoft.com/.default"}) + if err != nil { + return nil, fmt.Errorf("create graph service client: %w", err) + } + + return &Client{ + client: client, + }, nil +} + +func (c *Client) AddUserToGroup(ctx context.Context, groupId, userId string) error { + requestBody := models.NewReferenceCreate() + odataId := fmt.Sprintf("https://graph.microsoft.com/v1.0/directoryObjects/%s", userId) + requestBody.SetOdataId(&odataId) + if err := c.client.Groups().ByGroupId(groupId).Members().Ref().Post(ctx, requestBody, nil); err != nil { + return err + } + return nil +} + +func (c *Client) RemoveUserFromGroup(ctx context.Context, groupId, userId string) error { + if err := c.client.Groups().ByGroupId(groupId).Members().ByDirectoryObjectId(userId).Ref().Delete(ctx, nil); err != nil { + return err + } + return nil +} + +func (c *Client) CreateGroup(ctx context.Context, groupName string) (models.Groupable, error) { + requestBody := models.NewGroup() + requestBody.SetDisplayName(&groupName) + requestBody.SetSecurityEnabled(new(true)) + requestBody.SetMailEnabled(new(false)) + requestBody.SetMailNickname(&groupName) + requestBody.SetDescription(new("source:dapla-api")) + + group, err := c.client.Groups().Post(ctx, requestBody, nil) + if err != nil { + return nil, err + } + + return group, nil +} + +func (c *Client) GetGroup(ctx context.Context, groupId string) (models.Groupable, error) { + return c.client.Groups().ByGroupId(groupId).Get(ctx, nil) +} + +func (c *Client) GetTransitiveMembers(ctx context.Context, groupId string) ([]models.Userable, error) { + var users []models.Userable + req, err := c.client.Groups().ByGroupId(groupId).TransitiveMembers().GraphUser().Get(ctx, nil) + if err != nil { + return nil, fmt.Errorf("get entra id group members: %w", err) + } + + pageIterator, err := msgraphcore.NewPageIterator[models.Userable](req, c.client.GetAdapter(), models.CreateUserCollectionResponseFromDiscriminatorValue) + if err != nil { + return nil, fmt.Errorf("create entra id users pageiterator: %w", err) + } + + if err := pageIterator.Iterate(ctx, func(user models.Userable) bool { + users = append(users, user) + return true + }); err != nil { + return nil, fmt.Errorf("list all users group members: %w", err) + } + return users, nil +} + +func (c *Client) AssignAppRoleToGroup(ctx context.Context, groupId string, resourceId *uuid.UUID, appRoleId *uuid.UUID) error { + gcpSyncAssignment := models.NewAppRoleAssignment() + gcpSyncAssignment.SetAppRoleId(appRoleId) + gcpSyncAssignment.SetResourceId(resourceId) + groupUuid, err := uuid.Parse(groupId) + if err != nil { + return fmt.Errorf("parse group id: %w", err) + } + gcpSyncAssignment.SetPrincipalId(&groupUuid) + if _, err := c.client.Groups().ByGroupId(groupId).AppRoleAssignments().Post(ctx, gcpSyncAssignment, nil); err != nil { + return fmt.Errorf("create app role assignment: %w", err) + } + + return nil +} + +func (c *Client) GetAppRolesForGroup(ctx context.Context, groupId string) ([]models.AppRoleAssignmentable, error) { + appRolesResponse, err := c.client.Groups().ByGroupId(groupId).AppRoleAssignments().Get(ctx, nil) + if err != nil { + return nil, err + } + + pageIterator, _ := msgraphcore.NewPageIterator[models.AppRoleAssignmentable](appRolesResponse, c.client.GetAdapter(), models.CreateAppRoleAssignmentCollectionResponseFromDiscriminatorValue) + + var appRoleAssignments []models.AppRoleAssignmentable + if err := pageIterator.Iterate(ctx, func(apa models.AppRoleAssignmentable) bool { + appRoleAssignments = append(appRoleAssignments, apa) + return true + }); err != nil { + return nil, err + } + + return appRoleAssignments, nil +} diff --git a/internal/reconcilers/entraid/group/master/database/database.go b/internal/reconcilers/entraid/group/master/database/database.go index f249477..b0578e0 100644 --- a/internal/reconcilers/entraid/group/master/database/database.go +++ b/internal/reconcilers/entraid/group/master/database/database.go @@ -5,17 +5,20 @@ import ( "errors" "fmt" - msgraphsdk "github.com/microsoftgraph/msgraph-sdk-go" - "github.com/microsoftgraph/msgraph-sdk-go/models" "github.com/sirupsen/logrus" "github.com/statisticsnorway/dapla-api-reconcilers/internal/reconcilers/entraid/group/master" ) type Master struct { - client *msgraphsdk.GraphServiceClient + client entraIdClient } -func New(client *msgraphsdk.GraphServiceClient) *Master { +type entraIdClient interface { + AddUserToGroup(ctx context.Context, groupId, userId string) error + RemoveUserFromGroup(ctx context.Context, groupId string, userId string) error +} + +func New(client entraIdClient) *Master { return &Master{ client: client, } @@ -28,7 +31,7 @@ func (m *Master) Name() string { func (m *Master) RemoveUsers(ctx context.Context, group master.Group, localOnlyUsers, remoteOnlyUsers []master.User, log logrus.FieldLogger) error { var errs []error for _, user := range remoteOnlyUsers { - if err := m.client.Groups().ByGroupId(group.ExternalId).Members().ByDirectoryObjectId(user.ExternalId).Ref().Delete(ctx, nil); err != nil { + if err := m.client.RemoveUserFromGroup(ctx, group.ExternalId, user.ExternalId); err != nil { errs = append(errs, fmt.Errorf("remove user %q from group %q: %w", user.Email, group.Name, err)) } } @@ -38,10 +41,7 @@ func (m *Master) RemoveUsers(ctx context.Context, group master.Group, localOnlyU func (m *Master) AddUsers(ctx context.Context, group master.Group, localOnlyUsers, remoteOnlyUsers []master.User, log logrus.FieldLogger) error { var errs []error for _, user := range localOnlyUsers { - requestBody := models.NewReferenceCreate() - odataId := fmt.Sprintf("https://graph.microsoft.com/v1.0/directoryObjects/%s", user.ExternalId) - requestBody.SetOdataId(&odataId) - if err := m.client.Groups().ByGroupId(group.ExternalId).Members().Ref().Post(ctx, requestBody, nil); err != nil { + if err := m.client.AddUserToGroup(ctx, group.ExternalId, user.ExternalId); err != nil { errs = append(errs, fmt.Errorf("add user %q to group %q: %w", user.Email, group.Name, err)) } } diff --git a/internal/reconcilers/entraid/group/reconciler.go b/internal/reconcilers/entraid/group/reconciler.go index 445c232..f38761e 100644 --- a/internal/reconcilers/entraid/group/reconciler.go +++ b/internal/reconcilers/entraid/group/reconciler.go @@ -6,14 +6,11 @@ import ( "fmt" "slices" - "cloud.google.com/go/auth/credentials/idtoken" - "github.com/Azure/azure-sdk-for-go/sdk/azidentity" "github.com/google/uuid" - msgraphsdk "github.com/microsoftgraph/msgraph-sdk-go" - msgraphcore "github.com/microsoftgraph/msgraph-sdk-go-core" "github.com/microsoftgraph/msgraph-sdk-go/models" "github.com/microsoftgraph/msgraph-sdk-go/models/odataerrors" "github.com/sirupsen/logrus" + "github.com/statisticsnorway/dapla-api-reconcilers/internal/entraidclient" "github.com/statisticsnorway/dapla-api-reconcilers/internal/reconcilers" "github.com/statisticsnorway/dapla-api-reconcilers/internal/reconcilers/entraid/group/master" "github.com/statisticsnorway/dapla-api-reconcilers/internal/reconcilers/entraid/group/master/database" @@ -42,13 +39,24 @@ type syncQueuer interface { Add(group string, member *string) error } +type entraIdClient interface { + AddUserToGroup(ctx context.Context, groupId, userId string) error + RemoveUserFromGroup(ctx context.Context, groupId, userId string) error + CreateGroup(ctx context.Context, groupName string) (models.Groupable, error) + GetGroup(ctx context.Context, groupId string) (models.Groupable, error) + GetTransitiveMembers(ctx context.Context, groupId string) ([]models.Userable, error) + AssignAppRoleToGroup(ctx context.Context, groupId string, resourceId *uuid.UUID, appRoleId *uuid.UUID) error + GetAppRolesForGroup(ctx context.Context, groupId string) ([]models.AppRoleAssignmentable, error) +} + type entraIdGroupReconciler struct { - mainCtx context.Context - service *msgraphsdk.GraphServiceClient - entraIdConfig entraIdConfig - syncQueuer syncQueuer - masterHandler master.Handler - memberMasterConfig memberMasterConfig + mainCtx context.Context + entraIdClient entraIdClient + entraIdConfig entraIdConfig + staticEntraIdClient bool + syncQueuer syncQueuer + masterHandler master.Handler + memberMasterConfig memberMasterConfig } type entraIdConfig struct { @@ -65,12 +73,25 @@ type memberMasterConfig struct { overrides string } -func New(ctx context.Context, sq syncQueuer) reconcilers.Reconciler { +type OptFunc func(*entraIdGroupReconciler) + +func WithEntraIdClient(client entraIdClient) OptFunc { + return func(r *entraIdGroupReconciler) { + r.entraIdClient = client + r.staticEntraIdClient = true + } +} + +func New(ctx context.Context, sq syncQueuer, opts ...OptFunc) reconcilers.Reconciler { r := &entraIdGroupReconciler{ mainCtx: ctx, syncQueuer: sq, } + for _, opt := range opts { + opt(r) + } + return r } @@ -170,7 +191,7 @@ func (r *entraIdGroupReconciler) Reconcile(ctx context.Context, client *apiclien return nil } -func (r *entraIdGroupReconciler) reconcileGroup(ctx context.Context, entraId *msgraphsdk.GraphServiceClient, client *apiclient.APIClient, teamSlug string, groupName string, log logrus.FieldLogger) error { +func (r *entraIdGroupReconciler) reconcileGroup(ctx context.Context, entraId entraIdClient, client *apiclient.APIClient, teamSlug string, groupName string, log logrus.FieldLogger) error { log = log.WithField("groupName", groupName) group, created, err := getOrCreateGroup(ctx, entraId, client, groupName, r.entraIdConfig.GroupPrefix) if err != nil { @@ -201,7 +222,7 @@ func (r *entraIdGroupReconciler) reconcileGroup(ctx context.Context, entraId *ms return fmt.Errorf("get database members: %w", err) } - entraIdUsers, err := getEntraIdMembers(ctx, entraId, group.ExternalId) + entraIdUsers, err := entraId.GetTransitiveMembers(ctx, group.ExternalId) if err != nil { return fmt.Errorf("get entra id members: %w", err) } @@ -256,28 +277,7 @@ func getDatabaseMembers(ctx context.Context, client *apiclient.APIClient, group return dbMembers, dbMembersIt.Err() } -func getEntraIdMembers(ctx context.Context, entraId *msgraphsdk.GraphServiceClient, groupId string) ([]models.Userable, error) { - var entraIdUsers []models.Userable - entraIdUsersReq, err := entraId.Groups().ByGroupId(groupId).TransitiveMembers().GraphUser().Get(ctx, nil) - if err != nil { - return nil, fmt.Errorf("get entra id group members: %w", err) - } - - pageIterator, err := msgraphcore.NewPageIterator[models.Userable](entraIdUsersReq, entraId.GetAdapter(), models.CreateUserCollectionResponseFromDiscriminatorValue) - if err != nil { - return nil, fmt.Errorf("create entra id users pageiterator: %w", err) - } - - if err := pageIterator.Iterate(ctx, func(user models.Userable) bool { - entraIdUsers = append(entraIdUsers, user) - return true - }); err != nil { - return nil, fmt.Errorf("list all users group members: %w", err) - } - return entraIdUsers, nil -} - -func getOrCreateGroup(ctx context.Context, entraId *msgraphsdk.GraphServiceClient, client *apiclient.APIClient, groupName string, groupPrefix string) (_ *master.Group, created bool, err error) { +func getOrCreateGroup(ctx context.Context, entraId entraIdClient, client *apiclient.APIClient, groupName string, groupPrefix string) (_ *master.Group, created bool, err error) { dbGroup, err := client.Groups().Get(ctx, &protoapi.GetGroupRequest{ Name: groupName, }) @@ -286,8 +286,7 @@ func getOrCreateGroup(ctx context.Context, entraId *msgraphsdk.GraphServiceClien } entraIdGroupName := fmt.Sprintf("%s%s", groupPrefix, groupName) if dbGroup.Group.ExternalId != nil { - _, err := entraId.Groups().ByGroupId(*dbGroup.Group.ExternalId).Get(ctx, nil) - if err == nil { + if _, err := entraId.GetGroup(ctx, *dbGroup.Group.ExternalId); err == nil { return &master.Group{ ExternalId: *dbGroup.Group.ExternalId, Name: entraIdGroupName, @@ -301,14 +300,7 @@ func getOrCreateGroup(ctx context.Context, entraId *msgraphsdk.GraphServiceClien } } - requestBody := models.NewGroup() - requestBody.SetDisplayName(&entraIdGroupName) - requestBody.SetSecurityEnabled(new(true)) - requestBody.SetMailEnabled(new(false)) - requestBody.SetMailNickname(&entraIdGroupName) - requestBody.SetDescription(new("source:dapla-api")) - - group, err := entraId.Groups().Post(ctx, requestBody, nil) + group, err := entraId.CreateGroup(ctx, entraIdGroupName) if err != nil { return nil, false, fmt.Errorf("create group: %w", err) } @@ -320,27 +312,17 @@ func getOrCreateGroup(ctx context.Context, entraId *msgraphsdk.GraphServiceClien } // ensureAppRoles assigns any app roles which may be missing on the given group -func ensureAppRoles(ctx context.Context, entraId *msgraphsdk.GraphServiceClient, groupId string, appRoleId *uuid.UUID, resourceIds ...*uuid.UUID) error { - appRolesResponse, err := entraId.Groups().ByGroupId(groupId).AppRoleAssignments().Get(ctx, nil) +func ensureAppRoles(ctx context.Context, entraId entraIdClient, groupId string, appRoleId *uuid.UUID, resourceIds ...*uuid.UUID) error { + appRoleAssignments, err := entraId.GetAppRolesForGroup(ctx, groupId) if err != nil { return fmt.Errorf("get app role assignments for group %q: %w", groupId, err) } - pageIterator, _ := msgraphcore.NewPageIterator[models.AppRoleAssignmentable](appRolesResponse, entraId.GetAdapter(), models.CreateAppRoleAssignmentCollectionResponseFromDiscriminatorValue) - - var appRoleAssignments []models.AppRoleAssignmentable - if err := pageIterator.Iterate(ctx, func(apa models.AppRoleAssignmentable) bool { - appRoleAssignments = append(appRoleAssignments, apa) - return true - }); err != nil { - return fmt.Errorf("iterate through app role assignments for group %q: %w", groupId, err) - } - for _, resourceId := range resourceIds { if !slices.ContainsFunc(appRoleAssignments, func(apa models.AppRoleAssignmentable) bool { return apa.GetAppRoleId() != nil && *apa.GetAppRoleId() == *appRoleId && apa.GetResourceId() != nil && *apa.GetResourceId() == *resourceId }) { - if err := assignAppRole(ctx, entraId, groupId, resourceId, appRoleId); err != nil { + if err := entraId.AssignAppRoleToGroup(ctx, groupId, resourceId, appRoleId); err != nil { return fmt.Errorf("assign app role %q on resourceId %q for group %q: %w", appRoleId.String(), resourceId.String(), groupId, err) } } @@ -349,24 +331,6 @@ func ensureAppRoles(ctx context.Context, entraId *msgraphsdk.GraphServiceClient, return nil } -// assignAppRole sets the necessary App Role on the given Entra ID group, -// so that it can be synced by the GCP provisioning app. -func assignAppRole(ctx context.Context, entraId *msgraphsdk.GraphServiceClient, groupId string, resourceId *uuid.UUID, appRoleId *uuid.UUID) error { - gcpSyncAssignment := models.NewAppRoleAssignment() - gcpSyncAssignment.SetAppRoleId(appRoleId) - gcpSyncAssignment.SetResourceId(resourceId) - groupUuid, err := uuid.Parse(groupId) - if err != nil { - return fmt.Errorf("parse group id: %w", err) - } - gcpSyncAssignment.SetPrincipalId(&groupUuid) - if _, err := entraId.Groups().ByGroupId(groupId).AppRoleAssignments().Post(ctx, gcpSyncAssignment, nil); err != nil { - return fmt.Errorf("create app role assignment: %w", err) - } - - return nil -} - // getDatabaseOnlyUsers takes a list of database users and remote/Entra ID users and returns // those users which are only present in the database. These are the users that need to be added // to the Entra ID group. @@ -409,7 +373,7 @@ func getRemoteOnlyUsers(dbUsers []*protoapi.GroupMember, remoteUsers []models.Us return remoteOnly } -func (r *entraIdGroupReconciler) getEntraIdClient(config *protoapi.ConfigReconcilerResponse) (*msgraphsdk.GraphServiceClient, error) { +func (r *entraIdGroupReconciler) getEntraIdClient(config *protoapi.ConfigReconcilerResponse) (entraIdClient, error) { rc := entraIdConfig{} for _, c := range config.Nodes { switch c.Key { @@ -444,36 +408,21 @@ func (r *entraIdGroupReconciler) getEntraIdClient(config *protoapi.ConfigReconci } if rc == r.entraIdConfig { - return r.service, nil - } - - creds, err := azidentity.NewClientAssertionCredential(rc.TenantId, rc.ClientId, func(ctx context.Context) (string, error) { - creds, err := idtoken.NewCredentials(&idtoken.Options{Audience: "api://AzureADTokenExchange"}) - if err != nil { - return "", err - } - token, err := creds.Token(ctx) - if err != nil { - return "", err - } - return token.Value, nil - }, nil) - if err != nil { - return nil, fmt.Errorf("exchange for azure credentials: %w", err) + return r.entraIdClient, nil } - service, err := msgraphsdk.NewGraphServiceClientWithCredentials(creds, []string{"https://graph.microsoft.com/.default"}) + client, err := entraidclient.New(rc.TenantId, rc.ClientId) if err != nil { - return nil, fmt.Errorf("create graph service client: %w", err) + return nil, fmt.Errorf("create entraid client: %w", err) } - r.service = service + r.entraIdClient = client r.entraIdConfig = rc - return service, nil + return client, nil } -func (r *entraIdGroupReconciler) configureMemberMasters(apiClient *apiclient.APIClient, entraidClient *msgraphsdk.GraphServiceClient, config *protoapi.ConfigReconcilerResponse) error { +func (r *entraIdGroupReconciler) configureMemberMasters(apiClient *apiclient.APIClient, entraidClient entraIdClient, config *protoapi.ConfigReconcilerResponse) error { newConfig := memberMasterConfig{} for _, c := range config.Nodes { switch c.Key {