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"
|
p.ApprovalPrompt = "consent"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
return &AzureProvider{ProviderData: p}
|
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:
|
// Filters that will be possible to use:
|
||||||
// contains - unknown function | "https://graph.microsoft.com/v1.0/me/memberOf?$filter=contains(displayName,%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,%27olm%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%27olm%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"
|
requestUrl := "https://graph.microsoft.com/v1.0/me/memberOf?$select=displayName"
|
||||||
|
|
||||||
groups := make([]string, 0)
|
groups := make([]string, 0)
|
||||||
|
|
||||||
for {
|
for {
|
||||||
|
|
@ -166,6 +168,9 @@ func (p *AzureProvider) GetGroups(s *SessionState, f string) (string, error) {
|
||||||
req.Header.Add("Content-Type", "application/json")
|
req.Header.Add("Content-Type", "application/json")
|
||||||
|
|
||||||
groupData, err := api.Request(req)
|
groupData, err := api.Request(req)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
for _, groupInfo := range groupData.Get("value").MustArray() {
|
for _, groupInfo := range groupData.Get("value").MustArray() {
|
||||||
v, ok := groupInfo.(map[string]interface{})
|
v, ok := groupInfo.(map[string]interface{})
|
||||||
|
|
@ -218,6 +223,7 @@ func (p *AzureProvider) ValidateGroup(s *SessionState) bool {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
return false
|
||||||
}
|
}
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,13 +1,70 @@
|
||||||
package providers
|
package providers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"fmt"
|
||||||
"github.com/bmizerany/assert"
|
"github.com/bmizerany/assert"
|
||||||
|
"io/ioutil"
|
||||||
|
"log"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
"net/url"
|
"net/url"
|
||||||
|
"strings"
|
||||||
"testing"
|
"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 {
|
func testAzureProvider(hostname string) *AzureProvider {
|
||||||
p := NewAzureProvider(
|
p := NewAzureProvider(
|
||||||
&ProviderData{
|
&ProviderData{
|
||||||
|
|
@ -18,6 +75,7 @@ func testAzureProvider(hostname string) *AzureProvider {
|
||||||
ValidateURL: &url.URL{},
|
ValidateURL: &url.URL{},
|
||||||
ProtectedResource: &url.URL{},
|
ProtectedResource: &url.URL{},
|
||||||
Scope: ""})
|
Scope: ""})
|
||||||
|
|
||||||
if hostname != "" {
|
if hostname != "" {
|
||||||
updateURL(p.Data().LoginURL, hostname)
|
updateURL(p.Data().LoginURL, hostname)
|
||||||
updateURL(p.Data().RedeemURL, 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, "type assertion to string failed", err.Error())
|
||||||
assert.Equal(t, "", email)
|
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",
|
AccessToken: "token1234",
|
||||||
ExpiresOn: time.Now().Add(time.Duration(1) * time.Hour),
|
ExpiresOn: time.Now().Add(time.Duration(1) * time.Hour),
|
||||||
RefreshToken: "refresh4321",
|
RefreshToken: "refresh4321",
|
||||||
|
Groups: "test-group-1|test-group-2",
|
||||||
}
|
}
|
||||||
encoded, err := s.EncodeSessionState(c)
|
encoded, err := s.EncodeSessionState(c)
|
||||||
assert.Equal(t, nil, err)
|
assert.Equal(t, nil, err)
|
||||||
assert.Equal(t, 3, strings.Count(encoded, ":"))
|
assert.Equal(t, 4, strings.Count(encoded, ":"))
|
||||||
|
|
||||||
ss, err := DecodeSessionState(encoded, c)
|
ss, err := DecodeSessionState(encoded, c)
|
||||||
t.Logf("%#v", ss)
|
t.Logf("%#v", ss)
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue