Files
summercms/modules/bouncer/registry_test.go

309 lines
8.7 KiB
Go

package bouncer
import (
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
)
type writerGuard struct {
principal *Principal
err error
wrote *bool
}
func (g writerGuard) Authenticate(*http.Request) (*Principal, error) {
return g.principal, g.err
}
func (g writerGuard) WriteUnauthorized(w http.ResponseWriter, err error) {
if g.wrote != nil {
*g.wrote = true
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusUnauthorized)
_, _ = w.Write([]byte(`{"from":"writer"}`))
}
type credOnlyGuard struct {
principal *Principal
cred any
err error
}
func (g credOnlyGuard) AuthenticateCredential(*http.Request) (*Principal, any, error) {
return g.principal, g.cred, g.err
}
type notAGuard struct{}
func TestDuplicateGuardNameFailsWithPluginAndName(t *testing.T) {
reg := NewRegistry()
if err := reg.Register("golem15.user", "jwt", writerGuard{}); err != nil {
t.Fatal(err)
}
err := reg.Register("golem15.acme", "jwt", writerGuard{})
if err == nil || !strings.Contains(err.Error(), "golem15.user") || !strings.Contains(err.Error(), "jwt") {
t.Fatalf("want plugin and guard name in error, got %v", err)
}
}
func TestUnknownGuardNameFails(t *testing.T) {
reg := NewRegistry()
_, err := reg.Middleware("jwt")
if err == nil || !strings.Contains(err.Error(), "jwt") {
t.Fatalf("want guard name in error, got %v", err)
}
}
func TestRegisterNeitherInterfaceNamesPluginAndName(t *testing.T) {
reg := NewRegistry()
err := reg.Register("golem15.demo", "oops", notAGuard{})
if err == nil || !strings.Contains(err.Error(), "golem15.demo") || !strings.Contains(err.Error(), "oops") {
t.Fatalf("want plugin and guard name in error, got %v", err)
}
}
func TestWriterGuardFailureWritesOwnResponse(t *testing.T) {
wrote := false
reg := NewRegistry()
if err := reg.Register("golem15.user", "jwt", writerGuard{
err: errors.New("nope"),
wrote: &wrote,
}); err != nil {
t.Fatal(err)
}
mw, err := reg.Middleware("jwt")
if err != nil {
t.Fatal(err)
}
called := false
h := mw(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
called = true
}))
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil))
if !wrote {
t.Fatal("WriteUnauthorized was not called")
}
if called {
t.Fatal("next ran on guard failure")
}
if rec.Code != http.StatusUnauthorized || rec.Body.String() != `{"from":"writer"}` {
t.Fatalf("status=%d body=%q", rec.Code, rec.Body.String())
}
}
func TestCredentialGuardSoftFailAndSuccess(t *testing.T) {
reg := NewRegistry()
failing := credOnlyGuard{err: errors.New("bad token")}
if err := reg.Register("golem15.acme", "inv_token", failing); err != nil {
t.Fatal(err)
}
mw, err := reg.Middleware("inv_token")
if err != nil {
t.Fatal(err)
}
t.Run("failure-passes-through", func(t *testing.T) {
called := false
h := mw(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
called = true
if _, ok := User(r.Context()); ok {
t.Fatal("principal attached on failure")
}
if _, ok := Credential(r.Context()); ok {
t.Fatal("credential attached on failure")
}
w.WriteHeader(http.StatusNoContent)
}))
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil))
if !called {
t.Fatal("next did not run on CredentialGuard failure")
}
if rec.Code != http.StatusNoContent {
t.Fatalf("status = %d", rec.Code)
}
})
okReg := NewRegistry()
token := &struct{ ID uint }{ID: 9}
principal := &Principal{ID: 3}
if err := okReg.Register("golem15.acme", "inv_token", credOnlyGuard{
principal: principal,
cred: token,
}); err != nil {
t.Fatal(err)
}
okMW, err := okReg.Middleware("inv_token")
if err != nil {
t.Fatal(err)
}
t.Run("success-attaches-both", func(t *testing.T) {
h := okMW(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
got, ok := User(r.Context())
if !ok || got != principal {
t.Fatalf("user = %+v ok=%t", got, ok)
}
cred, ok := Credential(r.Context())
if !ok || cred != token {
t.Fatalf("cred = %#v ok=%t", cred, ok)
}
w.WriteHeader(http.StatusNoContent)
}))
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/", nil))
if rec.Code != http.StatusNoContent {
t.Fatalf("status = %d", rec.Code)
}
})
}
func TestJWTGuardRegistryMatchesMiddleware(t *testing.T) {
users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}}
direct := Middleware(secret, users)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Hit", "1")
w.WriteHeader(http.StatusNoContent)
}))
reg := NewRegistry()
if err := reg.Register("golem15.user", "jwt", NewJWTGuard(secret, users, nil)); err != nil {
t.Fatal(err)
}
viaReg, err := reg.Middleware("jwt")
if err != nil {
t.Fatal(err)
}
hReg := viaReg(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Hit", "1")
w.WriteHeader(http.StatusNoContent)
}))
expired := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(-time.Hour).Unix(),
}, []byte(secret))
badSig := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "1",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte("other-secret"))
unknown := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{
"sub": "99",
"exp": time.Now().Add(time.Hour).Unix(),
}, []byte(secret))
cases := []struct {
name string
header string
}{
{name: "missing"},
{name: "expired", header: "Bearer " + expired},
{name: "bad-signature", header: "Bearer " + badSig},
{name: "unknown-user", header: "Bearer " + unknown},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
req1 := httptest.NewRequest(http.MethodGet, "/", nil)
req2 := httptest.NewRequest(http.MethodGet, "/", nil)
if tc.header != "" {
req1.Header.Set("Authorization", tc.header)
req2.Header.Set("Authorization", tc.header)
}
rec1 := httptest.NewRecorder()
rec2 := httptest.NewRecorder()
direct.ServeHTTP(rec1, req1)
hReg.ServeHTTP(rec2, req2)
if rec1.Code != rec2.Code {
t.Fatalf("status direct=%d registry=%d", rec1.Code, rec2.Code)
}
if rec1.Body.String() != rec2.Body.String() {
t.Fatalf("body direct=%q registry=%q", rec1.Body.String(), rec2.Body.String())
}
if rec1.Header().Get("Content-Type") != rec2.Header().Get("Content-Type") {
t.Fatalf("content-type direct=%q registry=%q", rec1.Header().Get("Content-Type"), rec2.Header().Get("Content-Type"))
}
if rec1.Header().Get("X-Hit") != "" || rec2.Header().Get("X-Hit") != "" {
t.Fatal("handler ran")
}
})
}
}
func TestRegisterRejectsTypedNilGuard(t *testing.T) {
var reg Registry
var g *credPtrGuard
if err := reg.Register("p", "n", g); err == nil {
t.Fatal("typed-nil guard accepted")
}
}
type credPtrGuard struct{}
func (*credPtrGuard) Authenticate(*http.Request) (*Principal, error) { return nil, nil }
type nilFuncGuard func(*http.Request) (*Principal, error)
func (f nilFuncGuard) Authenticate(r *http.Request) (*Principal, error) { return f(r) }
type nilMapGuard map[string]string
func (nilMapGuard) Authenticate(*http.Request) (*Principal, error) { return nil, nil }
func TestRegisterRejectsTypedNilPointerFuncMapGuards(t *testing.T) {
var (
ptr *credPtrGuard
fn nilFuncGuard
mp nilMapGuard
)
for name, g := range map[string]any{"pointer": ptr, "func": fn, "map": mp} {
t.Run(name, func(t *testing.T) {
var reg Registry
err := reg.Register("golem15.p", "g", g)
if err == nil || !strings.Contains(err.Error(), "golem15.p") || !strings.Contains(err.Error(), `"g"`) {
t.Fatalf("typed-nil %s guard: err = %v", name, err)
}
if _, err := reg.Middleware("g"); err == nil {
t.Fatal("rejected guard must not be resolvable")
}
})
}
}
func TestRegisterAcceptsValidGuards(t *testing.T) {
var reg Registry
if err := reg.Register("p", "ptr", &credPtrGuard{}); err != nil {
t.Fatal(err)
}
if err := reg.Register("p", "fn", nilFuncGuard(func(*http.Request) (*Principal, error) { return nil, nil })); err != nil {
t.Fatal(err)
}
if err := reg.Register("p", "map", nilMapGuard{}); err != nil {
t.Fatal(err)
}
if err := reg.Register("p", "w", writerGuard{}); err != nil {
t.Fatal(err)
}
}
func TestRegistryOwner(t *testing.T) {
reg := NewRegistry()
if _, ok := reg.Owner("backend"); ok {
t.Fatal("an unregistered guard reported an owner")
}
if err := reg.Register("acme.owner", "backend", NewBackendJWTGuard("test-secret-for-registry-owner", nil, nil, nil)); err != nil {
t.Fatal(err)
}
if owner, ok := reg.Owner("backend"); !ok || owner != "acme.owner" {
t.Fatalf("Owner = %q, %v; want acme.owner", owner, ok)
}
var nilReg *Registry
if _, ok := nilReg.Owner("backend"); ok {
t.Fatal("a nil registry reported an owner")
}
}