Files
WorkBuddy/cnc-sales-backend/internal/auth/auth.go
T
2026-08-16 16:29:06 +08:00

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) + `}`))
}