295 lines
8.6 KiB
Go
295 lines
8.6 KiB
Go
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")
|
|
}
|
|
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)
|
|
}
|
|
}
|
|
}
|