package bouncer import ( "context" "encoding/json" "net/http" "net/http/httptest" "testing" "time" "github.com/golang-jwt/jwt/v5" ) const secret = "test-secret" type memUsers struct { byID map[uint]*Principal err error } func (m memUsers) FindByID(ctx context.Context, id uint) (*Principal, error) { if m.err != nil { return nil, m.err } return m.byID[id], nil } func sign(t *testing.T, method jwt.SigningMethod, claims jwt.MapClaims, key []byte) string { t.Helper() tok := jwt.NewWithClaims(method, claims) s, err := tok.SignedString(key) if err != nil { t.Fatal(err) } return s } func TestVerifyRejectsBadTokens(t *testing.T) { valid := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ "sub": "1", "exp": time.Now().Add(time.Hour).Unix(), }, []byte(secret)) if sub, err := Verify(valid, secret); err != nil || sub != "1" { t.Fatalf("valid token: %s %v", sub, err) } expired := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ "sub": "1", "exp": time.Now().Add(-time.Hour).Unix(), }, []byte(secret)) if _, err := Verify(expired, secret); err == nil || err.Error() != msgTokenExpired { t.Fatalf("expired: %v", err) } none := sign(t, jwt.SigningMethodHS384, jwt.MapClaims{ "sub": "1", "exp": time.Now().Add(time.Hour).Unix(), }, []byte(secret)) if _, err := Verify(none, secret); err == nil || err.Error() != msgBadSignature { t.Fatalf("wrong alg: %v", err) } badSig := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ "sub": "1", "exp": time.Now().Add(time.Hour).Unix(), }, []byte("other-secret")) if _, err := Verify(badSig, secret); err == nil || err.Error() != msgBadSignature { t.Fatalf("bad sig: %v", err) } if _, err := Verify("not-a-jwt", secret); err == nil || err.Error() != msgMalformed { t.Fatalf("malformed: %v", err) } missingSub := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ "exp": time.Now().Add(time.Hour).Unix(), }, []byte(secret)) if _, err := Verify(missingSub, secret); err == nil || err.Error() != msgRequiredClaims { t.Fatalf("missing sub: %v", err) } } func TestMiddlewareStatusBodies(t *testing.T) { users := memUsers{byID: map[uint]*Principal{1: {ID: 1}}} h := Middleware(secret, users)(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("X-Hit", "1") w.WriteHeader(http.StatusNoContent) })) assert401 := func(t *testing.T, req *http.Request, msg string) { t.Helper() rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusUnauthorized { t.Fatalf("status = %d body %s", rec.Code, rec.Body.String()) } if rec.Header().Get("X-Hit") != "" { t.Fatal("handler ran") } var body map[string]any if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { t.Fatal(err) } if body["error"] != true || body["message"] != msg { t.Fatalf("body = %v want %s", body, msg) } } t.Run("missing", func(t *testing.T) { assert401(t, httptest.NewRequest(http.MethodGet, "/", nil), msgTokenNotProvided) }) t.Run("unknown-user", func(t *testing.T) { tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ "sub": "99", "exp": time.Now().Add(time.Hour).Unix(), }, []byte(secret)) req := httptest.NewRequest(http.MethodGet, "/", nil) req.Header.Set("Authorization", "Bearer "+tok) assert401(t, req, msgUserNotFound) }) t.Run("valid", func(t *testing.T) { tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ "sub": "1", "exp": time.Now().Add(time.Hour).Unix(), }, []byte(secret)) req := httptest.NewRequest(http.MethodGet, "/", nil) req.Header.Set("Authorization", "Bearer "+tok) rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusNoContent || rec.Header().Get("X-Hit") != "1" { t.Fatalf("status=%d hit=%s", rec.Code, rec.Header().Get("X-Hit")) } }) }