Add unittests for group filtering (Azure) (#1)

* Add unittests for group filtering (Azure)
This commit is contained in:
Pavel Sorokin 2017-02-22 15:31:48 -08:00 committed by Brandon Matthews
parent 342cac58ca
commit 5b0e28fcbe
3 changed files with 164 additions and 4 deletions

View File

@ -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
}

View File

@ -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)
}

View File

@ -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)