Add unittests for group filtering (Azure) (#1)
* Add unittests for group filtering (Azure)
This commit is contained in:
parent
342cac58ca
commit
5b0e28fcbe
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in New Issue