package bouncer import ( "context" "encoding/base64" "encoding/json" "net/http" "net/http/httptest" "strings" "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") } if got := rec.Header().Get("Cache-Control"); got != "no-cache, private" { t.Fatalf("Cache-Control = %q", got) } 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")) } }) t.Run("malformed", func(t *testing.T) { req := httptest.NewRequest(http.MethodGet, "/", nil) req.Header.Set("Authorization", "Bearer not-a-jwt") assert401(t, req, msgMalformed) }) t.Run("alg-none", func(t *testing.T) { req := httptest.NewRequest(http.MethodGet, "/", nil) req.Header.Set("Authorization", "Bearer "+noneToken(t, jwt.MapClaims{ "sub": "1", "exp": time.Now().Add(time.Hour).Unix(), })) assert401(t, req, msgBadSignature) }) t.Run("absent-exp", func(t *testing.T) { tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": "1"}, []byte(secret)) req := httptest.NewRequest(http.MethodGet, "/", nil) req.Header.Set("Authorization", "Bearer "+tok) assert401(t, req, msgBadSignature) }) } func TestVerifyRejectsAlgNoneEmptySecretAndAbsentExp(t *testing.T) { valid := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ "sub": "1", "exp": time.Now().Add(time.Hour).Unix(), }, []byte(secret)) none := noneToken(t, jwt.MapClaims{ "sub": "1", "exp": time.Now().Add(time.Hour).Unix(), }) if _, err := Verify(none, secret); err == nil || err.Error() != msgBadSignature { t.Fatalf("alg none: %v", err) } if _, err := Verify(valid, ""); err == nil || !strings.Contains(err.Error(), "jwt secret is empty") { t.Fatalf("empty secret: %v", err) } if _, err := Verify(valid, " "); err == nil || !strings.Contains(err.Error(), "jwt secret is empty") { t.Fatalf("blank secret: %v", err) } noExp := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": "1"}, []byte(secret)) if sub, err := Verify(noExp, secret); err == nil || sub != "" { t.Fatalf("absent exp must fail, got %q %v", sub, err) } } func TestVerifyAndMiddlewareOmitTokenAndSecret(t *testing.T) { const leakSecret = "unique-hs256-secret-value-9f3a" tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{ "sub": "1", "exp": time.Now().Add(time.Hour).Unix(), }, []byte(leakSecret)) assertClean := func(t *testing.T, msg string) { t.Helper() if strings.Contains(msg, leakSecret) || strings.Contains(msg, tok) { t.Fatalf("leaked secret or token: %s", msg) } } if _, err := Verify("not-a-jwt", leakSecret); err == nil { t.Fatal("want malformed") } else { assertClean(t, err.Error()) } if _, err := Verify(tok, "other-"+leakSecret); err == nil { t.Fatal("want bad signature") } else { assertClean(t, err.Error()) } h := Middleware(leakSecret, memUsers{})(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { t.Fatal("handler must not run") })) req := httptest.NewRequest(http.MethodGet, "/", nil) req.Header.Set("Authorization", "Bearer "+tok) rec := httptest.NewRecorder() h.ServeHTTP(rec, req) if rec.Code != http.StatusUnauthorized { t.Fatalf("status = %d", rec.Code) } assertClean(t, rec.Body.String()) } func TestContextUserRoundTrip(t *testing.T) { if _, ok := User(t.Context()); ok { t.Fatal("empty context must have no user") } p := &Principal{ID: 7, MustChangePassword: true} got, ok := User(WithUser(t.Context(), p)) if !ok || got != p || got.ID != 7 || !got.MustChangePassword { t.Fatalf("got %+v ok=%t", got, ok) } } func noneToken(t *testing.T, claims jwt.MapClaims) string { t.Helper() payload, err := json.Marshal(claims) if err != nil { t.Fatal(err) } header := base64.RawURLEncoding.EncodeToString([]byte(`{"alg":"none","typ":"JWT"}`)) body := base64.RawURLEncoding.EncodeToString(payload) return header + "." + body + "." } func TestVerifyRejectsFractionalSubject(t *testing.T) { exp := time.Now().Add(time.Hour).Unix() tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": 12.5, "exp": exp}, []byte(secret)) if sub, err := Verify(tok, secret); err == nil || sub != "" { t.Fatalf("fractional sub accepted: %q %v", sub, err) } tok = sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": 12, "exp": exp}, []byte(secret)) if sub, err := Verify(tok, secret); err != nil || sub != "12" { t.Fatalf("whole sub: %q %v", sub, err) } } func TestVerifySubjectMatrix(t *testing.T) { exp := time.Now().Add(time.Hour).Unix() cases := []struct { name string sub any want string }{ {"fractional", 12.5, ""}, {"huge float", 1e300, ""}, {"2^60 float", float64(1 << 60), ""}, {"negative", -1, ""}, {"zero", 0, ""}, {"whole number", 12, "12"}, {"string", "12", "12"}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { tok := sign(t, jwt.SigningMethodHS256, jwt.MapClaims{"sub": tc.sub, "exp": exp}, []byte(secret)) sub, err := Verify(tok, secret) if tc.want == "" { if err == nil || sub != "" { t.Fatalf("sub %v accepted: %q", tc.sub, sub) } return } if err != nil || sub != tc.want { t.Fatalf("sub %v: %q %v", tc.sub, sub, err) } }) } } func TestSubjectJSONNumber(t *testing.T) { for in, want := range map[string]string{"1.5": "", "-3": "", "0": "", "12": "12"} { if got := subject(jwt.MapClaims{"sub": json.Number(in)}); got != want { t.Errorf("json.Number(%q) = %q, want %q", in, got, want) } } }