mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
adds user's transitiveMemberOf group lister to lib/msgraph (#58095)
* set and validate default graph endpoint * add IterateUsersTransitiveMemberOf method to list user groups * run make fix-imports * fix: return on group type error: - updates test
This commit is contained in:
+19
-7
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
]
|
||||
`
|
||||
|
||||
@@ -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/<user-id>/transitiveMemberOf
|
||||
// - Security groups: /v1.0/users/<user-id>/transitiveMemberOf/microsoft.graph.group?$filter=mailEnabled eq false and securityEnabled eq true
|
||||
// - Directory roles: /v1.0/users/<user-id>/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)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user