111 lines
3.0 KiB
Go
111 lines
3.0 KiB
Go
package auth
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"net/http"
|
|
"os"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/golang-jwt/jwt/v5"
|
|
"golang.org/x/crypto/bcrypt"
|
|
)
|
|
|
|
type ctxKey string
|
|
|
|
const (
|
|
keyUserID ctxKey = "userID"
|
|
keyIsAdmin ctxKey = "isAdmin"
|
|
)
|
|
|
|
// secret 返回 JWT 签名密钥,可用环境变量 JWT_SECRET 覆盖(生产必须设置)
|
|
func secret() []byte {
|
|
if s := os.Getenv("JWT_SECRET"); s != "" {
|
|
return []byte(s)
|
|
}
|
|
return []byte("dev-secret-change-me")
|
|
}
|
|
|
|
// HashPassword 生成 bcrypt 密码哈希
|
|
func HashPassword(pw string) (string, error) {
|
|
b, err := bcrypt.GenerateFromPassword([]byte(pw), bcrypt.DefaultCost)
|
|
return string(b), err
|
|
}
|
|
|
|
// CheckPassword 校验密码
|
|
func CheckPassword(hash, pw string) bool {
|
|
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(pw)) == nil
|
|
}
|
|
|
|
// GenerateToken 签发 72 小时有效的 JWT
|
|
func GenerateToken(userID int64) (string, error) {
|
|
claims := jwt.RegisteredClaims{
|
|
Subject: strconv.FormatInt(userID, 10),
|
|
IssuedAt: jwt.NewNumericDate(time.Now()),
|
|
ExpiresAt: jwt.NewNumericDate(time.Now().Add(72 * time.Hour)),
|
|
}
|
|
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(secret())
|
|
}
|
|
|
|
// Middleware JWT 鉴权中间件
|
|
func Middleware(next http.Handler) http.Handler {
|
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
h := r.Header.Get("Authorization")
|
|
if !strings.HasPrefix(h, "Bearer ") {
|
|
writeErr(w, http.StatusUnauthorized, "未登录")
|
|
return
|
|
}
|
|
token, err := jwt.ParseWithClaims(strings.TrimPrefix(h, "Bearer "), &jwt.RegisteredClaims{},
|
|
func(t *jwt.Token) (interface{}, error) {
|
|
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
|
|
return nil, errors.New("unexpected signing method")
|
|
}
|
|
return secret(), nil
|
|
})
|
|
if err != nil || !token.Valid {
|
|
writeErr(w, http.StatusUnauthorized, "登录已过期,请重新登录")
|
|
return
|
|
}
|
|
claims, ok := token.Claims.(*jwt.RegisteredClaims)
|
|
if !ok {
|
|
writeErr(w, http.StatusUnauthorized, "无效凭证")
|
|
return
|
|
}
|
|
uid, err := strconv.ParseInt(claims.Subject, 10, 64)
|
|
if err != nil || uid <= 0 {
|
|
writeErr(w, http.StatusUnauthorized, "无效凭证")
|
|
return
|
|
}
|
|
ctx := context.WithValue(r.Context(), keyUserID, uid)
|
|
next.ServeHTTP(w, r.WithContext(ctx))
|
|
})
|
|
}
|
|
|
|
// UserID 从请求上下文取当前用户 ID
|
|
func UserID(r *http.Request) int64 {
|
|
if v, ok := r.Context().Value(keyUserID).(int64); ok {
|
|
return v
|
|
}
|
|
return 0
|
|
}
|
|
|
|
// WithAdmin 标记当前用户为管理员(供业务层判断)
|
|
func WithAdmin(r *http.Request, isAdmin bool) *http.Request {
|
|
ctx := context.WithValue(r.Context(), keyIsAdmin, isAdmin)
|
|
return r.WithContext(ctx)
|
|
}
|
|
|
|
// IsAdmin 判断当前用户是否为管理员
|
|
func IsAdmin(r *http.Request) bool {
|
|
v, _ := r.Context().Value(keyIsAdmin).(bool)
|
|
return v
|
|
}
|
|
|
|
func writeErr(w http.ResponseWriter, code int, msg string) {
|
|
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
|
w.WriteHeader(code)
|
|
_, _ = w.Write([]byte(`{"error":` + strconv.Quote(msg) + `}`))
|
|
}
|