diff --git a/token/jwt.go b/token/jwt.go index ffc3e6a..fd6db4a 100644 --- a/token/jwt.go +++ b/token/jwt.go @@ -188,6 +188,22 @@ func parseJWT(token string, keys []map[string]any) (*JWT, error) { return (*JWT)(t), nil } +func matchGroup(pth, group string, includeSubgroups bool) bool { + if !includeSubgroups { + return pth == "/group/"+group+"/" + } + + if !strings.HasPrefix(pth, "/group/") { + return false + } + + if !strings.HasSuffix(pth, "/") { + return false + } + + return strings.HasPrefix("/group/"+group+"/", pth) +} + func (token *JWT) Check(host, group string, username *string) (string, []string, error) { sub, err := token.Claims.GetSubject() if err != nil { @@ -205,7 +221,7 @@ func (token *JWT) Check(host, group string, username *string) (string, []string, if !ok { return "", nil, errors.New("unexpected type for token") } - incSubgroups, _ := claims["include-subgroups"].(bool) + includeSubgroups, _ := claims["include-subgroups"].(bool) ok = false for _, u := range aud { @@ -221,13 +237,7 @@ func (token *JWT) Check(host, group string, username *string) (string, []string, continue } } - // aud path takes the form /group// - if !strings.HasPrefix(url.Path, "/group/") || !strings.HasSuffix(url.Path, "/") { - continue - } - tokenGroup := url.Path[len("/group/") : len(url.Path)-1] - if group == tokenGroup || - incSubgroups && (tokenGroup == "" || strings.HasPrefix(group, tokenGroup+"/")) { + if matchGroup(url.Path, group, includeSubgroups) { ok = true break } diff --git a/token/jwt_test.go b/token/jwt_test.go index 48bc781..eb091ef 100644 --- a/token/jwt_test.go +++ b/token/jwt_test.go @@ -77,6 +77,42 @@ func TestJWKRS256(t *testing.T) { } } +func TestMatchGroup(t *testing.T) { + type tt struct { + p, g string + s bool + } + good := []tt{ + {"/group/a/", "a", false}, + {"/group/a/b/", "a/b", false}, + {"/group/a/", "a", true}, + {"/group/a/", "a/b", true}, + {"/group/a/b/", "a/b", true}, + {"/group/a/b/", "a/b/c", true}, + } + + bad := []tt{ + {"/group/a/", "b", false}, + {"/group/a", "a", false}, + {"/group/a", "a", true}, + {"/group/a/", "a/b", false}, + {"/group/a/b/", "a", false}, + {"/group/a/b/", "a", true}, + } + + for _, test := range good { + if !matchGroup(test.p, test.g, test.s) { + t.Errorf("%v %v %v didn't match", test.p, test.g, test.s) + } + } + + for _, test := range bad { + if matchGroup(test.p, test.g, test.s) { + t.Errorf("%v %v %v matched", test.p, test.g, test.s) + } + } +} + func TestJWT(t *testing.T) { key := `{"alg":"HS256","k":"H7pCkktUl5KyPCZ7CKw09y1j460tfIv4dRcS1XstUKY","key_ops":["sign","verify"],"kty":"oct"}` var k map[string]interface{}