Files
helix-proxy/internal/auth/jwt.go
T
s1d3sw1ped_bot 22844a2a67
Format / gofmt (push) Successful in 8s
Format / gofmt (pull_request) Successful in 9s
CI / Build (push) Successful in 24s
CI / Build (pull_request) Successful in 25s
CI / Go Tests (pull_request) Successful in 39s
CI / Go Tests (push) Successful in 40s
Persist a per-install JWT signing secret instead of a compiled-in default.
Admin tokens were forgeable whenever JWT_SECRET was unset. Prefer the env var, otherwise write a random key to data/.jwt_secret.
2026-08-31 23:55:34 +00:00

181 lines
5.0 KiB
Go

package auth
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"net/http"
"os"
"path/filepath"
"strings"
"time"
"github.com/golang-jwt/jwt/v5"
)
var (
ErrInvalidToken = errors.New("invalid or expired token")
ErrNoToken = errors.New("no authorization token")
)
// Claims for our JWT (minimal, matching original style).
type Claims struct {
UserID int `json:"user_id"`
Email string `json:"email"`
Name string `json:"name"`
Roles []string `json:"roles"`
jwt.RegisteredClaims
}
// Context key for user info.
type contextKey string
const UserContextKey contextKey = "user"
// JWTManager handles signing and validation.
type JWTManager struct {
secret []byte
}
const jwtSecretFilename = ".jwt_secret"
// LoadOrCreateSecret returns JWT_SECRET from the environment if set, otherwise
// a per-install secret persisted at dir/.jwt_secret (created on first run).
func LoadOrCreateSecret(dir string) (string, error) {
if s := strings.TrimSpace(os.Getenv("JWT_SECRET")); s != "" {
return s, nil
}
if strings.TrimSpace(dir) == "" {
return "", fmt.Errorf("jwt secret directory is required when JWT_SECRET is unset")
}
if err := os.MkdirAll(dir, 0o755); err != nil {
return "", fmt.Errorf("create jwt secret dir: %w", err)
}
path := filepath.Join(dir, jwtSecretFilename)
if b, err := os.ReadFile(path); err == nil {
s := strings.TrimSpace(string(b))
if s != "" {
return s, nil
}
} else if !errors.Is(err, os.ErrNotExist) {
return "", fmt.Errorf("read jwt secret: %w", err)
}
s, err := randomSecret()
if err != nil {
return "", err
}
if err := os.WriteFile(path, []byte(s+"\n"), 0o600); err != nil {
return "", fmt.Errorf("write jwt secret: %w", err)
}
return s, nil
}
func randomSecret() (string, error) {
b := make([]byte, 32)
if _, err := rand.Read(b); err != nil {
return "", fmt.Errorf("generate jwt secret: %w", err)
}
return hex.EncodeToString(b), nil
}
// NewJWTManager creates a manager. Pass a secret from LoadOrCreateSecret or JWT_SECRET.
// An empty secret is replaced with a random in-memory value (tokens will not survive restart).
func NewJWTManager(secret string) *JWTManager {
if secret == "" {
s, err := randomSecret()
if err != nil {
panic(err)
}
secret = s
}
return &JWTManager{secret: []byte(secret)}
}
// GenerateToken creates a JWT for the given user (1 hour expiry like typical).
func (m *JWTManager) GenerateToken(userID int, email, name string, roles []string) (string, error) {
now := time.Now()
if roles == nil {
roles = []string{}
}
claims := Claims{
UserID: userID,
Email: email,
Name: name,
Roles: roles,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(now.Add(1 * time.Hour)),
IssuedAt: jwt.NewNumericDate(now),
NotBefore: jwt.NewNumericDate(now),
Issuer: "helix-proxy",
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
return token.SignedString(m.secret)
}
// ValidateToken parses and validates the token, returns claims.
func (m *JWTManager) ValidateToken(tokenStr string) (*Claims, error) {
token, err := jwt.ParseWithClaims(tokenStr, &Claims{}, func(t *jwt.Token) (interface{}, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, fmt.Errorf("unexpected signing method: %v", t.Header["alg"])
}
return m.secret, nil
})
if err != nil {
return nil, ErrInvalidToken
}
if claims, ok := token.Claims.(*Claims); ok && token.Valid {
return claims, nil
}
return nil, ErrInvalidToken
}
// Middleware returns a chi/http middleware that validates Bearer token and injects user into context.
// Skips if no token (for public endpoints like /login we handle separately).
// On failure returns 401.
func (m *JWTManager) Middleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
authHeader := r.Header.Get("Authorization")
if authHeader == "" {
// No token: let handler decide (some endpoints public)
next.ServeHTTP(w, r)
return
}
parts := strings.SplitN(authHeader, " ", 2)
if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") {
http.Error(w, "invalid authorization header", http.StatusUnauthorized)
return
}
claims, err := m.ValidateToken(parts[1])
if err != nil {
http.Error(w, "invalid token", http.StatusUnauthorized)
return
}
// Inject into context
ctx := context.WithValue(r.Context(), UserContextKey, claims)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
// GetUserFromContext extracts claims if present (for handlers that require auth).
func GetUserFromContext(ctx context.Context) (*Claims, bool) {
claims, ok := ctx.Value(UserContextKey).(*Claims)
return claims, ok
}
// RequireAuth is a simple helper that can be used inside handlers for protected routes
// (alternative to middleware if you want per-route).
func RequireAuth(w http.ResponseWriter, r *http.Request) (*Claims, bool) {
claims, ok := GetUserFromContext(r.Context())
if !ok {
http.Error(w, "authentication required", http.StatusUnauthorized)
return nil, false
}
return claims, true
}