feat(03-01): add ServeMux groups, JWT verifier, and serve command
Named middleware resolves at boot, HS256 tokens are pinned with required exp/sub, and both binaries expose a signal-aware serve command. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
28
bouncer/context.go
Normal file
28
bouncer/context.go
Normal file
@@ -0,0 +1,28 @@
|
||||
package bouncer
|
||||
|
||||
import "context"
|
||||
|
||||
type userKey struct{}
|
||||
|
||||
// Principal is the authenticated identity stored on the request context.
|
||||
type Principal struct {
|
||||
ID uint
|
||||
MustChangePassword bool
|
||||
}
|
||||
|
||||
// WithUser stores the verified principal on ctx.
|
||||
func WithUser(ctx context.Context, user *Principal) context.Context {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
return context.WithValue(ctx, userKey{}, user)
|
||||
}
|
||||
|
||||
// User returns the verified principal from ctx.
|
||||
func User(ctx context.Context) (*Principal, bool) {
|
||||
if ctx == nil {
|
||||
return nil, false
|
||||
}
|
||||
u, ok := ctx.Value(userKey{}).(*Principal)
|
||||
return u, ok && u != nil
|
||||
}
|
||||
139
bouncer/jwt.go
Normal file
139
bouncer/jwt.go
Normal file
@@ -0,0 +1,139 @@
|
||||
package bouncer
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
const (
|
||||
msgTokenNotProvided = "Token not provided"
|
||||
msgTokenExpired = "Token has expired"
|
||||
msgUserNotFound = "User not found"
|
||||
msgBadSignature = "Token Signature could not be verified."
|
||||
msgMalformed = "Wrong number of segments"
|
||||
msgRequiredClaims = "JWT payload does not contain the required claims"
|
||||
)
|
||||
|
||||
// UserProvider loads a persisted user by JWT subject.
|
||||
type UserProvider interface {
|
||||
FindByID(ctx context.Context, id uint) (*Principal, error)
|
||||
}
|
||||
|
||||
// Middleware validates a pinned HS256 bearer token and loads the user.
|
||||
func Middleware(secret string, users UserProvider) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
raw, err := bearerToken(r)
|
||||
if err != nil {
|
||||
write401(w, err.Error())
|
||||
return
|
||||
}
|
||||
sub, err := Verify(raw, secret)
|
||||
if err != nil {
|
||||
write401(w, err.Error())
|
||||
return
|
||||
}
|
||||
id, err := strconv.ParseUint(sub, 10, 64)
|
||||
if err != nil || id == 0 {
|
||||
write401(w, msgUserNotFound)
|
||||
return
|
||||
}
|
||||
if users == nil {
|
||||
write401(w, msgUserNotFound)
|
||||
return
|
||||
}
|
||||
user, err := users.FindByID(r.Context(), uint(id))
|
||||
if err != nil {
|
||||
write401(w, "Authentication error")
|
||||
return
|
||||
}
|
||||
if user == nil {
|
||||
write401(w, msgUserNotFound)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r.WithContext(WithUser(r.Context(), user)))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Verify parses a token with HS256 pinned and a required exp and sub.
|
||||
func Verify(tokenString, secret string) (string, error) {
|
||||
if strings.TrimSpace(secret) == "" {
|
||||
return "", fmt.Errorf("bouncer: jwt secret is empty")
|
||||
}
|
||||
parser := jwt.NewParser(jwt.WithValidMethods([]string{"HS256"}), jwt.WithExpirationRequired())
|
||||
claims := jwt.MapClaims{}
|
||||
_, err := parser.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (any, error) {
|
||||
return []byte(secret), nil
|
||||
})
|
||||
if err != nil {
|
||||
return "", mapJWTError(err)
|
||||
}
|
||||
sub := subject(claims)
|
||||
if sub == "" {
|
||||
return "", errors.New(msgRequiredClaims)
|
||||
}
|
||||
return sub, nil
|
||||
}
|
||||
|
||||
func bearerToken(r *http.Request) (string, error) {
|
||||
h := strings.TrimSpace(r.Header.Get("Authorization"))
|
||||
if h == "" {
|
||||
return "", errors.New(msgTokenNotProvided)
|
||||
}
|
||||
token, ok := strings.CutPrefix(h, "Bearer ")
|
||||
token = strings.TrimSpace(token)
|
||||
if !ok || token == "" {
|
||||
return "", errors.New(msgTokenNotProvided)
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func subject(claims jwt.MapClaims) string {
|
||||
switch v := claims["sub"].(type) {
|
||||
case string:
|
||||
return strings.TrimSpace(v)
|
||||
case float64:
|
||||
if v <= 0 {
|
||||
return ""
|
||||
}
|
||||
return strconv.FormatInt(int64(v), 10)
|
||||
case json.Number:
|
||||
return strings.TrimSpace(v.String())
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func mapJWTError(err error) error {
|
||||
switch {
|
||||
case errors.Is(err, jwt.ErrTokenExpired):
|
||||
return errors.New(msgTokenExpired)
|
||||
case errors.Is(err, jwt.ErrTokenSignatureInvalid), errors.Is(err, jwt.ErrTokenUnverifiable):
|
||||
return errors.New(msgBadSignature)
|
||||
case errors.Is(err, jwt.ErrTokenMalformed):
|
||||
return errors.New(msgMalformed)
|
||||
default:
|
||||
msg := err.Error()
|
||||
if strings.Contains(strings.ToLower(msg), "expired") {
|
||||
return errors.New(msgTokenExpired)
|
||||
}
|
||||
if strings.Contains(strings.ToLower(msg), "malformed") || strings.Contains(msg, "segment") {
|
||||
return errors.New(msgMalformed)
|
||||
}
|
||||
return errors.New(msgBadSignature)
|
||||
}
|
||||
}
|
||||
|
||||
func write401(w http.ResponseWriter, message string) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{"error": true, "message": message})
|
||||
}
|
||||
134
bouncer/jwt_test.go
Normal file
134
bouncer/jwt_test.go
Normal file
@@ -0,0 +1,134 @@
|
||||
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"))
|
||||
}
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user