diff --git a/providers/azure.go b/providers/azure.go index e114b5ec..c131ef3f 100644 --- a/providers/azure.go +++ b/providers/azure.go @@ -41,6 +41,7 @@ func NewAzureProvider(p *ProviderData) *AzureProvider { p.ApprovalPrompt = "consent" } + return &AzureProvider{ProviderData: p} } @@ -149,11 +150,12 @@ func (p *AzureProvider) GetGroups(s *SessionState, f string) (string, error) { */ // // Filters that will be possible to use: - // contains - unknown function | "https://graph.microsoft.com/v1.0/me/memberOf?$filter=contains(displayName,%27olm%27)" - // startswith - not supported | "https://graph.microsoft.com/v1.0/me/memberOf?$filter=startswith(displayName,%27olm%27)" - // substring - not supported | "https://graph.microsoft.com/v1.0/me/memberOf?$filter=substring(displayName,0,2)%20eq%20%27olm%27" + // contains - unknown function | "https://graph.microsoft.com/v1.0/me/memberOf?$filter=contains(displayName,%27groupname%27)" + // startswith - not supported | "https://graph.microsoft.com/v1.0/me/memberOf?$filter=startswith(displayName,%27groupname%27)" + // substring - not supported | "https://graph.microsoft.com/v1.0/me/memberOf?$filter=substring(displayName,0,2)%20eq%20%27groupname%27" requestUrl := "https://graph.microsoft.com/v1.0/me/memberOf?$select=displayName" + groups := make([]string, 0) for { @@ -166,6 +168,9 @@ func (p *AzureProvider) GetGroups(s *SessionState, f string) (string, error) { req.Header.Add("Content-Type", "application/json") groupData, err := api.Request(req) + if err != nil { + return "", err + } for _, groupInfo := range groupData.Get("value").MustArray() { v, ok := groupInfo.(map[string]interface{}) @@ -218,6 +223,7 @@ func (p *AzureProvider) ValidateGroup(s *SessionState) bool { return true } } + return false } return true } diff --git a/providers/azure_test.go b/providers/azure_test.go index 22d6ccba..c8102901 100644 --- a/providers/azure_test.go +++ b/providers/azure_test.go @@ -1,13 +1,70 @@ package providers import ( + "fmt" "github.com/bmizerany/assert" + "io/ioutil" + "log" "net/http" "net/http/httptest" "net/url" + "strings" "testing" ) +var ( + path_group string = "/v1.0/me/memberOf?$select=displayName" + path_group_next string = "/v1.0/me/memberOf?$select=displayName&$skiptoken=X%27test-token%27" + path_group_wrong string = "/v1.0/him/memberOf?$select=displayName" + payload_group_empty string = `{"@odata.context":"https://graph.microsoft.com/v1.0/$metadata#directoryObjects(displayName)","value":[]}` + payload_group_garbage string = `{"@odata.context":"https://graph.microsoft.com/v1.0/$metadata#directoryObjects(displayName)","value":[{"@odata.type":"#microsoft.graph.group","displayName":"test-group-1"},{"@odata.type":"#microsoft.graph.group","displayName":"test-group-2"}]}` + payload_group_simple string = `{"@odata.context":"https://graph.microsoft.com/v1.0/$metadata#directoryObjects(displayName)","value":[{"@odata.type":"#microsoft.graph.group","displayName":"test-group-1"},{"@odata.type":"#microsoft.graph.group","displayName":"test-group-2"}]}` + payload_group_part_1 string = `{"@odata.context":"https://graph.microsoft.com/v1.0/$metadata#directoryObjects(displayName)","@odata.nextLink":"https://graph.microsoft.com/v1.0/me/memberOf?$select=displayName&$skiptoken=X%27test-token%27","value":[{"@odata.type":"#microsoft.graph.group","displayName":"test-group-1"},{"@odata.type":"#microsoft.graph.group","displayName":"test-group-2"}]}` + payload_group_part_2 string = `{"@odata.context":"https://graph.microsoft.com/v1.0/$metadata#directoryObjects(displayName)","value":[{"@odata.type":"#microsoft.graph.group","displayName":"test-group-3"}]}` +) + +type mockTransport struct { + params map[string]string +} + +func (t *mockTransport) RoundTrip(req *http.Request) (*http.Response, error) { + log.Printf("Starting Round Tripper") + // Create mocked http.Response + response := &http.Response{ + Header: make(http.Header), + Request: req, + StatusCode: http.StatusOK, + } + response.Header.Set("Content-Type", "application/json") + + //url := req.URL + full_request := req.URL.Path + if req.URL.RawQuery != "" { + full_request += "?" + req.URL.RawQuery + } + var err error + if value, ok := t.params[full_request]; ok { + if req.Header.Get("Authorization") != "Bearer imaginary_access_token" { + response.StatusCode = http.StatusForbidden + err = fmt.Errorf("got 403. Bearer token '%v' is not correct", req.Header.Get("Authorization")) + } else { + response.StatusCode = http.StatusOK + response.Body = ioutil.NopCloser(strings.NewReader(value)) + err = nil + } + + } else { + response.StatusCode = http.StatusNotFound + err = fmt.Errorf("got 404. Requested path '%v' is not found", full_request) + } + + return response, err +} + +func newMockTransport(params map[string]string) http.RoundTripper { + return &mockTransport{params} +} + func testAzureProvider(hostname string) *AzureProvider { p := NewAzureProvider( &ProviderData{ @@ -18,6 +75,7 @@ func testAzureProvider(hostname string) *AzureProvider { ValidateURL: &url.URL{}, ProtectedResource: &url.URL{}, Scope: ""}) + if hostname != "" { updateURL(p.Data().LoginURL, hostname) updateURL(p.Data().RedeemURL, hostname) @@ -198,3 +256,98 @@ func TestAzureProviderGetEmailAddressIncorrectOtherMails(t *testing.T) { assert.Equal(t, "type assertion to string failed", err.Error()) assert.Equal(t, "", email) } + +func TestAzureProviderNoGroups(t *testing.T) { + params := map[string]string{ + path_group: payload_group_empty} + + http.DefaultClient.Transport = newMockTransport(params) + + p := testAzureProvider("") + + session := &SessionState{ + AccessToken: "imaginary_access_token", + IDToken: "imaginary_IDToken_token"} + + groups, err := p.GetGroups(session, "") + http.DefaultClient.Transport = nil + + assert.Equal(t, nil, err) + assert.Equal(t, "", groups) +} + +func TestAzureProviderWrongRequestGroups(t *testing.T) { + params := map[string]string{ + path_group_wrong: payload_group_part_1} + http.DefaultClient.Transport = newMockTransport(params) + log.Printf("Def %#v\n\n", http.DefaultClient.Transport) + + p := testAzureProvider("") + + session := &SessionState{ + AccessToken: "imaginary_access_token", + IDToken: "imaginary_IDToken_token"} + + groups, err := p.GetGroups(session, "") + http.DefaultClient.Transport = nil + + assert.NotEqual(t, nil, err) + assert.Equal(t, "", groups) +} + +func TestAzureProviderMultiRequestsGroups(t *testing.T) { + params := map[string]string{ + path_group: payload_group_part_1, + path_group_next: payload_group_part_2} + http.DefaultClient.Transport = newMockTransport(params) + + p := testAzureProvider("") + + session := &SessionState{ + AccessToken: "imaginary_access_token", + IDToken: "imaginary_IDToken_token"} + + groups, err := p.GetGroups(session, "") + http.DefaultClient.Transport = nil + + assert.Equal(t, nil, err) + assert.Equal(t, "test-group-1|test-group-2|test-group-3", groups) +} + +func TestAzureEmptyPermittedGroups(t *testing.T) { + p := testAzureProvider("") + + session := &SessionState{ + AccessToken: "imaginary_access_token", + IDToken: "imaginary_IDToken_token", + Groups: "no one|cares|non existing|groups"} + result := p.ValidateGroup(session) + + assert.Equal(t, true, result) +} + +func TestAzureWrongPermittedGroups(t *testing.T) { + p := testAzureProvider("") + p.SetGroupRestriction([]string{"test-group-2"}) + + session := &SessionState{ + AccessToken: "imaginary_access_token", + IDToken: "imaginary_IDToken_token", + Groups: "no one|cares|non existing|groups|test-group-1"} + result := p.ValidateGroup(session) + + assert.Equal(t, false, result) +} + +func TestAzureRightPermittedGroups(t *testing.T) { + p := testAzureProvider("") + p.SetGroupRestriction([]string{"test-group-1", "test-group-2"}) + + session := &SessionState{ + AccessToken: "imaginary_access_token", + IDToken: "imaginary_IDToken_token", + Groups: "no one|cares|test-group-2|non existing|groups"} + result := p.ValidateGroup(session) + + assert.Equal(t, true, result) +} diff --git a/providers/session_state_test.go b/providers/session_state_test.go index 6447f999..af689d0e 100644 --- a/providers/session_state_test.go +++ b/providers/session_state_test.go @@ -22,10 +22,11 @@ func TestSessionStateSerialization(t *testing.T) { AccessToken: "token1234", ExpiresOn: time.Now().Add(time.Duration(1) * time.Hour), RefreshToken: "refresh4321", + Groups: "test-group-1|test-group-2", } encoded, err := s.EncodeSessionState(c) assert.Equal(t, nil, err) - assert.Equal(t, 3, strings.Count(encoded, ":")) + assert.Equal(t, 4, strings.Count(encoded, ":")) ss, err := DecodeSessionState(encoded, c) t.Logf("%#v", ss)