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) } }