248 lines
6.9 KiB
Go
248 lines
6.9 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.fonoteka", "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.fonoteka", "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.fonoteka", "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)); 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 }
|