This commit is contained in:
@@ -0,0 +1,227 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
func TestBeginAuthorizationUsesStateNonceAndPKCE(t *testing.T) {
|
||||
now := time.Date(2026, 8, 1, 12, 0, 0, 0, time.UTC)
|
||||
flow, err := BeginAuthorization(oauth2.Endpoint{AuthURL: "https://auth.example/authorize"}, OIDCConfig{ClientID: "pulse", RedirectURL: "https://pulse.example/callback"}, now)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
parsed, err := url.Parse(flow.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
query := parsed.Query()
|
||||
for _, key := range []string{"state", "nonce", "code_challenge", "code_challenge_method"} {
|
||||
if query.Get(key) == "" {
|
||||
t.Fatalf("authorization URL missing %s", key)
|
||||
}
|
||||
}
|
||||
if query.Get("state") != flow.State || query.Get("nonce") != flow.Nonce || query.Get("code_challenge_method") != "S256" {
|
||||
t.Fatalf("authorization URL does not match flow: %s", flow.URL)
|
||||
}
|
||||
if query.Get("code_challenge") != pkceChallenge(flow.CodeVerifier) {
|
||||
t.Fatal("authorization URL has incorrect PKCE challenge")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateCallbackRejectsStateNonceFlowAbuse(t *testing.T) {
|
||||
now := time.Now()
|
||||
flow := Authorization{State: "expected", Nonce: "nonce", CodeVerifier: "verifier", ExpiresAt: now.Add(time.Minute)}
|
||||
if err := ValidateCallback(flow, "wrong", "code", now); err == nil {
|
||||
t.Fatal("wrong state was accepted")
|
||||
}
|
||||
if err := ValidateCallback(flow, flow.State, "", now); err == nil {
|
||||
t.Fatal("empty code was accepted")
|
||||
}
|
||||
flow.ExpiresAt = now.Add(-time.Second)
|
||||
if err := ValidateCallback(flow, flow.State, "code", now); err == nil {
|
||||
t.Fatal("expired flow was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRoleMappingAndAuthorizationMatrix(t *testing.T) {
|
||||
mapping := map[string]Role{"pulse-view": RoleViewer, "pulse-operator": RoleOperator, "pulse-admin": RoleAdministrator}
|
||||
role, err := MapRoles([]string{"unrelated", "pulse-operator", "pulse-view"}, mapping)
|
||||
if err != nil || role != RoleOperator {
|
||||
t.Fatalf("role mapping = %q, %v", role, err)
|
||||
}
|
||||
if _, err := MapRoles([]string{"unrelated"}, mapping); err == nil {
|
||||
t.Fatal("unmapped claims were authorized")
|
||||
}
|
||||
for _, test := range []struct {
|
||||
role Role
|
||||
permission Permission
|
||||
allowed bool
|
||||
}{
|
||||
{RoleViewer, PermissionView, true}, {RoleViewer, PermissionEdit, false},
|
||||
{RoleOperator, PermissionOperate, true}, {RoleOperator, PermissionAdmin, false},
|
||||
{RoleEditor, PermissionEdit, true}, {RoleEditor, PermissionAdmin, false},
|
||||
{RoleAdministrator, PermissionAdmin, true},
|
||||
} {
|
||||
if got := Allows(test.role, test.permission); got != test.allowed {
|
||||
t.Errorf("Allows(%s, %s) = %v, want %v", test.role, test.permission, got, test.allowed)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnauthorizedPathsAreDenied(t *testing.T) {
|
||||
handler := Require(PermissionEdit, http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) { response.WriteHeader(http.StatusNoContent) }))
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
ctx context.Context
|
||||
status int
|
||||
}{
|
||||
{"anonymous", context.Background(), http.StatusUnauthorized},
|
||||
{"viewer", WithPrincipal(context.Background(), Principal{Subject: "user-1", Role: RoleViewer}), http.StatusForbidden},
|
||||
{"editor", WithPrincipal(context.Background(), Principal{Subject: "user-1", Role: RoleEditor}), http.StatusNoContent},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
request := httptest.NewRequest(http.MethodGet, "/protected", nil).WithContext(test.ctx)
|
||||
response := httptest.NewRecorder()
|
||||
handler.ServeHTTP(response, request)
|
||||
if response.Code != test.status {
|
||||
t.Fatalf("status = %d, want %d", response.Code, test.status)
|
||||
}
|
||||
if test.status == http.StatusUnauthorized || test.status == http.StatusForbidden {
|
||||
if response.Header().Get("Cache-Control") != "private, no-store" {
|
||||
t.Fatalf("cache control = %q", response.Header().Get("Cache-Control"))
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBreakGlassIsDisabledByDefault(t *testing.T) {
|
||||
if (BreakGlassPolicy{}).Allows() {
|
||||
t.Fatal("break-glass unexpectedly enabled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverResolvesEndpointAndVerifier(t *testing.T) {
|
||||
var issuer string
|
||||
server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) {
|
||||
if request.URL.Path != "/.well-known/openid-configuration" {
|
||||
response.WriteHeader(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
response.Header().Set("Content-Type", "application/json")
|
||||
_, _ = response.Write([]byte(`{"issuer":"` + issuer + `","authorization_endpoint":"` + issuer + `/authorize","token_endpoint":"` + issuer + `/token","jwks_uri":"` + issuer + `/jwks","id_token_signing_alg_values_supported":["RS256"]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
issuer = server.URL
|
||||
|
||||
discovery, err := Discover(context.Background(), OIDCConfig{Issuer: issuer, ClientID: "pulse"})
|
||||
if err != nil {
|
||||
t.Fatalf("Discover: %v", err)
|
||||
}
|
||||
if discovery.Endpoint.AuthURL != issuer+"/authorize" || discovery.Endpoint.TokenURL != issuer+"/token" || discovery.Verifier == nil {
|
||||
t.Fatalf("discovery = %#v", discovery.Endpoint)
|
||||
}
|
||||
if _, err := NewVerifier(context.Background(), OIDCConfig{Issuer: issuer, ClientID: "pulse"}); err != nil {
|
||||
t.Fatalf("NewVerifier: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoverRejectsIncompleteOrUnreachableIssuer(t *testing.T) {
|
||||
unreachable := httptest.NewServer(http.NewServeMux())
|
||||
unreachable.Close()
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
config OIDCConfig
|
||||
}{
|
||||
{"missing issuer", OIDCConfig{ClientID: "pulse"}},
|
||||
{"missing client id", OIDCConfig{Issuer: "https://idp.example"}},
|
||||
{"unreachable issuer", OIDCConfig{Issuer: unreachable.URL, ClientID: "pulse"}},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
discovery, err := Discover(context.Background(), test.config)
|
||||
if err == nil {
|
||||
t.Fatal("incomplete configuration was accepted")
|
||||
}
|
||||
if discovery.Verifier != nil {
|
||||
t.Fatal("a verifier was returned with an error")
|
||||
}
|
||||
if strings.Contains(err.Error(), test.config.Issuer) && test.config.Issuer != "" {
|
||||
t.Fatalf("error leaks the issuer: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractIdentityRequiresToken(t *testing.T) {
|
||||
if _, err := ExtractIdentity(nil, "groups"); err == nil {
|
||||
t.Fatal("nil token was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeGroupClaimBoundsAndShapes(t *testing.T) {
|
||||
many, err := json.Marshal(make([]string, maxIdentityGroups+50))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
raw string
|
||||
want []string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "list", raw: `["pulse-admin"," pulse-view ",""]`, want: []string{"pulse-admin", "pulse-view"}},
|
||||
{name: "single string", raw: `"pulse-admin"`, want: []string{"pulse-admin"}},
|
||||
{name: "empty list", raw: `[]`, want: []string{}},
|
||||
{name: "object", raw: `{"groups":["pulse-admin"]}`, wantErr: true},
|
||||
{name: "number", raw: `7`, wantErr: true},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
groups, err := normalizeGroupClaim(json.RawMessage(test.raw))
|
||||
if (err != nil) != test.wantErr {
|
||||
t.Fatalf("err = %v, wantErr = %v", err, test.wantErr)
|
||||
}
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if len(groups) != len(test.want) {
|
||||
t.Fatalf("groups = %#v, want %#v", groups, test.want)
|
||||
}
|
||||
for index, value := range test.want {
|
||||
if groups[index] != value {
|
||||
t.Fatalf("groups = %#v, want %#v", groups, test.want)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
bounded, err := normalizeGroupClaim(many)
|
||||
if err != nil {
|
||||
t.Fatalf("normalizeGroupClaim: %v", err)
|
||||
}
|
||||
if len(bounded) != 0 {
|
||||
t.Fatalf("blank group values were kept: %d", len(bounded))
|
||||
}
|
||||
filled := make([]string, maxIdentityGroups+50)
|
||||
for index := range filled {
|
||||
filled[index] = "group"
|
||||
}
|
||||
encoded, err := json.Marshal(filled)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
capped, err := normalizeGroupClaim(encoded)
|
||||
if err != nil {
|
||||
t.Fatalf("normalizeGroupClaim: %v", err)
|
||||
}
|
||||
if len(capped) != maxIdentityGroups {
|
||||
t.Fatalf("groups = %d, want %d", len(capped), maxIdentityGroups)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user