diff --git a/api/types/msgraph.go b/api/types/msgraph.go index 59060c2bf67..e01ee323c20 100644 --- a/api/types/msgraph.go +++ b/api/types/msgraph.go @@ -55,3 +55,20 @@ func ValidateMSGraphEndpoints(loginEndpoint, graphEndpoint string) error { return nil } + +const ( + // EntraIDSecurityGroups represents security enabled Entra ID groups. + EntraIDSecurityGroups = "security-groups" + // EntraIDDirectoryRoles represents Entra ID directory roles. + EntraIDDirectoryRoles = "directory-roles" + // EntraIDAllGroups represents all types of Entra ID groups, including directory roles. + EntraIDAllGroups = "all-groups" +) + +// EntraIDGroupsTypes defines supported Entra ID +// group types for Entra ID groups proivder. +var EntraIDGroupsTypes = []string{ + EntraIDSecurityGroups, + EntraIDDirectoryRoles, + EntraIDAllGroups, +} diff --git a/lib/msgraph/client.go b/lib/msgraph/client.go index b074bc1f4c8..0c62419e1ef 100644 --- a/lib/msgraph/client.go +++ b/lib/msgraph/client.go @@ -35,12 +35,13 @@ import ( "github.com/jonboulle/clockwork" apidefaults "github.com/gravitational/teleport/api/defaults" + "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/api/utils/retryutils" "github.com/gravitational/teleport/lib/defaults" ) -// baseURL is the default value for [client.baseURL]. It is the address of MS Graph API v1.0. -const baseURL = "https://graph.microsoft.com/v1.0" +// graphVersion is the default version of the MS Graph API endpoint. +const graphVersion = "v1.0" // defaultPageSize is the page size used when [Config.PageSize] is not specified. const defaultPageSize = 500 @@ -85,6 +86,8 @@ type Config struct { RetryConfig *retryutils.RetryV2Config // PageSize limits the number of objects to return in one batch when using paginated requests (via the `$top` parameter). PageSize int + // GraphEndpoint specifies root domain of the Graph API. + GraphEndpoint string } // SetDefaults sets the default values for optional fields. @@ -101,6 +104,9 @@ func (cfg *Config) SetDefaults() { if cfg.PageSize <= 0 { cfg.PageSize = defaultPageSize } + if cfg.GraphEndpoint == "" { + cfg.GraphEndpoint = types.MSGraphDefaultEndpoint + } } // Validate checks that required fields are set. @@ -111,6 +117,9 @@ func (cfg *Config) Validate() error { if cfg.HTTPClient == nil { return trace.BadParameter("HTTPClient must be set") } + if err := types.ValidateMSGraphEndpoints("", cfg.GraphEndpoint); err != nil { + return trace.Wrap(err) + } return nil } @@ -129,7 +138,7 @@ func NewClient(cfg Config) (*Client, error) { if err := cfg.Validate(); err != nil { return nil, trace.Wrap(err) } - uri, err := url.Parse(baseURL) + base, err := url.Parse(cfg.GraphEndpoint) if err != nil { return nil, trace.Wrap(err) } @@ -138,14 +147,14 @@ func NewClient(cfg Config) (*Client, error) { tokenProvider: cfg.TokenProvider, clock: cfg.Clock, retryConfig: *cfg.RetryConfig, - baseURL: uri, + baseURL: base.JoinPath(graphVersion), pageSize: cfg.PageSize, }, nil } // request is the base function for HTTP API calls. // It implements retry handling in case of API throttling, see [https://learn.microsoft.com/en-us/graph/throttling]. -func (c *Client) request(ctx context.Context, method string, uri string, payload []byte) (*http.Response, error) { +func (c *Client) request(ctx context.Context, method string, uri string, header map[string]string, payload []byte) (*http.Response, error) { var body io.ReadSeeker = nil if len(payload) > 0 { body = bytes.NewReader(payload) @@ -166,6 +175,9 @@ func (c *Client) request(ctx context.Context, method string, uri string, payload return nil, trace.Wrap(err, "failed to get azure authentication token") } req.Header.Add("Authorization", "Bearer "+token.Token) + for i := range header { + req.Header.Add(i, header[i]) + } const maxRetries = 5 var retryAfter time.Duration @@ -255,7 +267,7 @@ func roundtrip[T any](ctx context.Context, c *Client, method string, uri string, return zero, trace.Wrap(err) } } - resp, err := c.request(ctx, method, uri, body) + resp, err := c.request(ctx, method, uri, nil /* extra headers */, body) if err != nil { return zero, trace.Wrap(err) } @@ -275,7 +287,7 @@ func (c *Client) patch(ctx context.Context, uri string, in any) error { if err != nil { return trace.Wrap(err) } - resp, err := c.request(ctx, http.MethodPatch, uri, body) + resp, err := c.request(ctx, http.MethodPatch, uri, nil /* extra headers */, body) if err != nil { return trace.Wrap(err) } diff --git a/lib/msgraph/client_test.go b/lib/msgraph/client_test.go index 174b8f924ce..59baed88320 100644 --- a/lib/msgraph/client_test.go +++ b/lib/msgraph/client_test.go @@ -36,6 +36,7 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/api/utils/retryutils" ) @@ -557,3 +558,179 @@ func TestGetApplication(t *testing.T) { } func toPtr[T any](s T) *T { return &s } + +func TestNewClient(t *testing.T) { + tests := []struct { + name string + config Config + expectedGraphEndpoint string + errExpected bool + errAssertion require.ErrorAssertionFunc + }{ + { + name: "empty endpoint sets default graph endpoint", + config: Config{ + TokenProvider: &fakeTokenProvider{}, + GraphEndpoint: "", + }, + expectedGraphEndpoint: types.MSGraphDefaultEndpoint, + errAssertion: require.NoError, + }, + { + name: "configured endpoint", + config: Config{ + TokenProvider: &fakeTokenProvider{}, + GraphEndpoint: "https://dod-graph.microsoft.us", + }, + expectedGraphEndpoint: "https://dod-graph.microsoft.us", + errAssertion: require.NoError, + }, + { + name: "invalid endpoint", + config: Config{ + TokenProvider: &fakeTokenProvider{}, + GraphEndpoint: "https://graph.windows.net", + }, + errExpected: true, + errAssertion: require.Error, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + clt, err := NewClient(test.config) + test.errAssertion(t, err) + if !test.errExpected { + require.Equal(t, test.expectedGraphEndpoint+"/"+graphVersion, clt.baseURL.String()) + } + }) + } +} + +func TestIterateUsersTransitiveMemberOf(t *testing.T) { + userID := "9ef1bc41-1b26-4a66-b8bc-956b2a54f8dc" + allGroupsPath := fmt.Sprintf("/%s/users/%s/transitiveMemberOf", graphVersion, userID) + groupsPath := fmt.Sprintf("/%s/users/%s/transitiveMemberOf/%s", graphVersion, userID, graphNamespaceGroups) + directoryRolePath := fmt.Sprintf("/%s/users/%s/transitiveMemberOf/%s", graphVersion, userID, graphNamespaceDirectoryRoles) + + consistencyHeaderValue := "" + foundQuery := make(url.Values) + requestedPath := "" + withRequestChecker := func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestedPath = r.URL.Path + consistencyHeaderValue = r.Header.Get("ConsistencyLevel") + foundQuery = r.URL.Query() + next.ServeHTTP(w, r) + }) + } + + mux := http.NewServeMux() + var groups []json.RawMessage + require.NoError(t, json.Unmarshal([]byte(userGroups), &groups)) + mux.Handle("GET "+allGroupsPath, withRequestChecker(paginatedHandler(t, groups))) + mux.Handle("GET "+groupsPath, withRequestChecker(paginatedHandler(t, groups))) + mux.Handle("GET "+directoryRolePath, withRequestChecker(paginatedHandler(t, groups))) + srv := httptest.NewServer(mux) + t.Cleanup(func() { srv.Close() }) + + uri, err := url.Parse(srv.URL) + require.NoError(t, err) + client := &Client{ + httpClient: &http.Client{}, + tokenProvider: &fakeTokenProvider{}, + retryConfig: retryConfig, + baseURL: uri.JoinPath(graphVersion), + pageSize: 2, // smaller page size so we actually fetch multiple pages with our small test payload + } + + assertConsistencyLevelHeader := func(t *testing.T, h string) { + t.Helper() + require.Equal(t, "eventual", h, "request made without ConsistencyLevel header") + } + assertCountQuery := func(t *testing.T, c string) { + t.Helper() + require.Equal(t, "true", c, "request made without $count query") + } + assertRequestedPath := func(t *testing.T, e, p string) { + t.Helper() + require.Equal(t, e, p, "expected request path did not match") + } + + t.Run(types.EntraIDSecurityGroups, func(tt *testing.T) { + var groupIDs []string + err := client.IterateUsersTransitiveMemberOf(context.Background(), userID, types.EntraIDSecurityGroups, func(group *Group) bool { + groupIDs = append(groupIDs, *group.ID) + return true + }) + require.NoError(t, err) + require.Len(t, groupIDs, 5) + + filterValue, err := url.QueryUnescape(foundQuery.Get("$filter")) + require.NoError(t, err) + require.Equal(t, securityGroupsFilter, filterValue, "security groups request made without filter query") + assertConsistencyLevelHeader(t, consistencyHeaderValue) + assertRequestedPath(t, groupsPath, requestedPath) + assertCountQuery(t, foundQuery.Get("$count")) + }) + + t.Run(types.EntraIDAllGroups, func(tt *testing.T) { + var groupIDs []string + err := client.IterateUsersTransitiveMemberOf(context.Background(), userID, types.EntraIDAllGroups, func(group *Group) bool { + groupIDs = append(groupIDs, *group.ID) + return true + }) + require.NoError(t, err) + require.Len(t, groupIDs, 5) + + require.Empty(t, foundQuery.Get("$filter"), "non security groups request made with filter query") + assertConsistencyLevelHeader(t, consistencyHeaderValue) + assertRequestedPath(t, allGroupsPath, requestedPath) + assertCountQuery(t, foundQuery.Get("$count")) + }) + + t.Run(types.EntraIDDirectoryRoles, func(tt *testing.T) { + var groupIDs []string + err := client.IterateUsersTransitiveMemberOf(context.Background(), userID, types.EntraIDDirectoryRoles, func(group *Group) bool { + groupIDs = append(groupIDs, *group.ID) + return true + }) + require.NoError(t, err) + require.Len(t, groupIDs, 5) + + require.Empty(t, foundQuery.Get("$filter"), "non security groups request made with filter query") + assertConsistencyLevelHeader(t, consistencyHeaderValue) + assertRequestedPath(t, directoryRolePath, requestedPath) + assertCountQuery(t, foundQuery.Get("$count")) + }) + + t.Run("unsupported-group-type", func(tt *testing.T) { + var groupIDs []string + err := client.IterateUsersTransitiveMemberOf(context.Background(), userID, "unsupported-group-type", func(group *Group) bool { + groupIDs = append(groupIDs, *group.ID) + return true + }) + require.Error(t, err) + }) +} + +var userGroups = ` +[ + { + "id": "07af5ddc-0f6b-4237-8b3c-64815501d1d5" + }, + { + "id": "dd034a93-4ac3-4095-8b9e-f521ad7eace9" + }, + { + "id": "20b81a2c-fda0-41e7-8268-48a014be0b08" + }, + { + "id": "97336101-e9a4-4455-9d19-945fd9178ff6" + }, + { + "id": "76c1db72-be9c-4ed5-8a42-bdeec6a34502" + } +] +` diff --git a/lib/msgraph/paginated.go b/lib/msgraph/paginated.go index a0b9488af9d..28b990b31af 100644 --- a/lib/msgraph/paginated.go +++ b/lib/msgraph/paginated.go @@ -27,12 +27,14 @@ import ( "strconv" "github.com/gravitational/trace" + + "github.com/gravitational/teleport/api/types" ) // iterateSimple implements pagination for "simple" object lists, where additional logic isn't needed func iterateSimple[T any](c *Client, ctx context.Context, endpoint string, f func(*T) bool) error { var err error - itErr := c.iterate(ctx, endpoint, func(msg json.RawMessage) bool { + itErr := c.iterate(ctx, endpoint, nil /* query */, nil /* optional header */, func(msg json.RawMessage) bool { var page []T if err = json.Unmarshal(msg, &page); err != nil { return false @@ -51,13 +53,18 @@ func iterateSimple[T any](c *Client, ctx context.Context, endpoint string, f fun } // iterate implements pagination for "list" endpoints. -func (c *Client) iterate(ctx context.Context, endpoint string, f func(json.RawMessage) bool) error { +func (c *Client) iterate(ctx context.Context, endpoint string, query url.Values, header map[string]string, f func(json.RawMessage) bool) error { uri := *c.baseURL uri.Path = path.Join(uri.Path, endpoint) - uri.RawQuery = url.Values{"$top": {strconv.Itoa(c.pageSize)}}.Encode() + pageSize := strconv.Itoa(c.pageSize) + if query == nil { + query = make(url.Values) + } + query.Add("$top", pageSize) + uri.RawQuery = query.Encode() uriString := uri.String() for uriString != "" { - resp, err := c.request(ctx, http.MethodGet, uriString, nil) + resp, err := c.request(ctx, http.MethodGet, uriString, header, nil /* payload */) if err != nil { return trace.Wrap(err) } @@ -115,7 +122,7 @@ func (c *Client) IterateServicePrincipals(ctx context.Context, f func(principal // Ref: [https://learn.microsoft.com/en-us/graph/api/group-list-members]. func (c *Client) IterateGroupMembers(ctx context.Context, groupID string, f func(GroupMember) bool) error { var err error - itErr := c.iterate(ctx, path.Join("groups", groupID, "members"), func(msg json.RawMessage) bool { + itErr := c.iterate(ctx, path.Join("groups", groupID, "members"), nil /* query */, nil /* optional header */, func(msg json.RawMessage) bool { var page []json.RawMessage if err = json.Unmarshal(msg, &page); err != nil { return false @@ -144,3 +151,66 @@ func (c *Client) IterateGroupMembers(ctx context.Context, groupID string, f func } return trace.Wrap(itErr) } + +const ( + graphNamespaceGroups = "microsoft.graph.group" + graphNamespaceDirectoryRoles = "microsoft.graph.directoryRole" +) + +const ( + securityGroupsFilter = `mailEnabled eq false and securityEnabled eq true` +) + +// IterateUsersTransitiveMemberOf lists groups that the user is a member of +// through a direct or nested group membership. +// This method calls user's transitiveMemberOf endpoint https://learn.microsoft.com/en-us/graph/api/user-list-transitivememberof?view=graph-rest-1.0&tabs=http. +// Supported endpoints: +// - All groups and directory roles: /v1.0/users//transitiveMemberOf +// - Security groups: /v1.0/users//transitiveMemberOf/microsoft.graph.group?$filter=mailEnabled eq false and securityEnabled eq true +// - Directory roles: /v1.0/users//transitiveMemberOf/microsoft.graph.directoryRole +// Only group ID is extracted from the response, so the DirectoryObject struct is sufficient +// to parse groups as well ass directory roles response. +func (c *Client) IterateUsersTransitiveMemberOf(ctx context.Context, userID, groupType string, f func(*Group) bool) error { + // MS Graph expects $count query parameter and + // "ConsistencyLevel: eventual" header set when using + // advanced query parameter such as $filter. + // https://learn.microsoft.com/en-us/graph/aad-advanced-queries?tabs=http#legend + query := url.Values{ + "$select": {"id"}, + "$count": {"true"}, + } + header := map[string]string{ + "ConsistencyLevel": "eventual", + } + + endpoint := path.Join("users", userID, "transitiveMemberOf") + switch groupType { + case types.EntraIDAllGroups: + // default endpoint suffices. + case types.EntraIDDirectoryRoles: + endpoint = path.Join(endpoint, graphNamespaceDirectoryRoles) + case types.EntraIDSecurityGroups: + endpoint = path.Join(endpoint, graphNamespaceGroups) + query.Add("$filter", securityGroupsFilter) + default: + return trace.BadParameter("unexpected group type %q received, expected types are %q", groupType, types.EntraIDGroupsTypes) + } + + var err error + itErr := c.iterate(ctx, endpoint, query, header, func(msg json.RawMessage) bool { + var page []Group + if err = json.Unmarshal(msg, &page); err != nil { + return false + } + for _, item := range page { + if !f(&item) { + return false + } + } + return true + }) + if err != nil { + return trace.Wrap(err) + } + return trace.Wrap(itErr) +}