做统一认证登录
This commit is contained in:
+321
-63
@@ -1,63 +1,321 @@
|
||||
package jwtutil
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
// 密钥(后续可从配置中读取)
|
||||
var secret = []byte("yunzer_jwt_secret_key")
|
||||
|
||||
// Claims 定义JWT的claims结构
|
||||
type Claims struct {
|
||||
UserID int `json:"user_id"`
|
||||
Username string `json:"username"`
|
||||
TenantId int `json:"tenant_id"` // 租户ID
|
||||
UserType string `json:"user_type"` // 用户类型:"user" / "employee" / "platform" 等
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// GenerateToken 生成JWT token
|
||||
func GenerateToken(userID int, username string, tenantId int, userType string) (string, error) {
|
||||
expirationTime := time.Now().Add(24 * time.Hour)
|
||||
|
||||
claims := &Claims{
|
||||
UserID: userID,
|
||||
Username: username,
|
||||
TenantId: tenantId,
|
||||
UserType: userType,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(expirationTime),
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
NotBefore: jwt.NewNumericDate(time.Now()),
|
||||
},
|
||||
}
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
tokenString, err := token.SignedString(secret)
|
||||
return tokenString, err
|
||||
}
|
||||
|
||||
// ParseToken 解析JWT token
|
||||
func ParseToken(tokenString string) (*Claims, error) {
|
||||
claims := &Claims{}
|
||||
token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (interface{}, error) {
|
||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, errors.New("unexpected signing method")
|
||||
}
|
||||
return secret, nil
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if !token.Valid {
|
||||
return nil, errors.New("invalid token")
|
||||
}
|
||||
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
package jwtutil
|
||||
|
||||
import (
|
||||
"crypto/rsa"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"math/big"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
beego "github.com/beego/beego/v2/server/web"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
// 算法常量
|
||||
const (
|
||||
AlgHS256 = "HS256"
|
||||
AlgRS256 = "RS256"
|
||||
)
|
||||
|
||||
// DefaultKid 默认密钥 ID。历史 token 未携带 kid,统一按此 ID 处理。
|
||||
const DefaultKid = "default"
|
||||
|
||||
// legacySecret 兼容用的内置密钥:仅用于解析历史已签发的 token。
|
||||
// 生产环境必须在 app.conf 中配置 jwt_secret,否则会一直使用该固定值。
|
||||
const legacySecret = "yunzer_jwt_secret_key"
|
||||
|
||||
// Claims JWT 载荷。旧字段保持不变,新增字段均为 omitempty,
|
||||
// 历史 token 解析后为零值,不影响现有 78 处解析点。
|
||||
type Claims struct {
|
||||
UserID int `json:"user_id"`
|
||||
Username string `json:"username"`
|
||||
TenantId int `json:"tenant_id"` // 租户ID
|
||||
UserType string `json:"user_type"` // 用户类型:"user" / "employee" / "platform" 等
|
||||
|
||||
// ---- 统一认证中心(OIDC)扩展字段 ----
|
||||
ClientID string `json:"client_id,omitempty"` // 接入应用 client_id,即 aud 的业务标识
|
||||
Sid string `json:"sid,omitempty"` // 会话ID,用于单点登出/踢下线
|
||||
Scope string `json:"scope,omitempty"`
|
||||
Amr string `json:"amr,omitempty"` // 认证方式:pwd/sms/otp/wx/...
|
||||
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 密钥管理
|
||||
|
||||
var (
|
||||
keyOnce sync.Once
|
||||
hsKeys map[string][]byte // kid -> HMAC 密钥
|
||||
rsaPriv *rsa.PrivateKey // RS256 签名
|
||||
rsaPub *rsa.PublicKey // RS256 验签
|
||||
rsaKid string // RS256 密钥 ID
|
||||
issuerVal string
|
||||
)
|
||||
|
||||
func loadKeys() {
|
||||
hsKeys = map[string][]byte{DefaultKid: []byte(legacySecret)}
|
||||
|
||||
// 主密钥(覆盖内置默认值)
|
||||
if s, _ := beego.AppConfig.String("jwt_secret"); strings.TrimSpace(s) != "" {
|
||||
hsKeys[DefaultKid] = []byte(strings.TrimSpace(s))
|
||||
}
|
||||
// 轮换密钥:jwt_secrets = kid1:secret1,kid2:secret2
|
||||
if s, _ := beego.AppConfig.String("jwt_secrets"); strings.TrimSpace(s) != "" {
|
||||
for _, item := range strings.Split(s, ",") {
|
||||
kv := strings.SplitN(strings.TrimSpace(item), ":", 2)
|
||||
if len(kv) == 2 && strings.TrimSpace(kv[0]) != "" && strings.TrimSpace(kv[1]) != "" {
|
||||
hsKeys[strings.TrimSpace(kv[0])] = []byte(strings.TrimSpace(kv[1]))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RS256 密钥对:优先读文件路径,其次读内联 PEM
|
||||
privPEM := readKeyConf("jwt_rsa_private_key", "jwt_rsa_private_key_file")
|
||||
pubPEM := readKeyConf("jwt_rsa_public_key", "jwt_rsa_public_key_file")
|
||||
if privPEM != "" {
|
||||
if k, err := jwt.ParseRSAPrivateKeyFromPEM([]byte(privPEM)); err == nil {
|
||||
rsaPriv = k
|
||||
rsaPub = &k.PublicKey
|
||||
}
|
||||
}
|
||||
if pubPEM != "" {
|
||||
if k, err := jwt.ParseRSAPublicKeyFromPEM([]byte(pubPEM)); err == nil {
|
||||
rsaPub = k
|
||||
}
|
||||
}
|
||||
if rsaKid == "" {
|
||||
rsaKid, _ = beego.AppConfig.String("jwt_rsa_kid")
|
||||
}
|
||||
if rsaKid == "" {
|
||||
rsaKid = "rsa-1"
|
||||
}
|
||||
issuerVal, _ = beego.AppConfig.String("jwt_issuer")
|
||||
}
|
||||
|
||||
func readKeyConf(inlineKey, fileKey string) string {
|
||||
if v, _ := beego.AppConfig.String(inlineKey); strings.TrimSpace(v) != "" {
|
||||
return strings.ReplaceAll(strings.TrimSpace(v), `\n`, "\n")
|
||||
}
|
||||
if p, _ := beego.AppConfig.String(fileKey); strings.TrimSpace(p) != "" {
|
||||
if b, err := os.ReadFile(strings.TrimSpace(p)); err == nil {
|
||||
return string(b)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func ensureKeys() {
|
||||
keyOnce.Do(loadKeys)
|
||||
}
|
||||
|
||||
// keyFunc 按 token header 的 alg + kid 选择验签密钥。
|
||||
// 未携带 kid 或 kid 未注册时回落到默认密钥,保证历史 token 仍可解析。
|
||||
func keyFunc(token *jwt.Token) (interface{}, error) {
|
||||
ensureKeys()
|
||||
kid, _ := token.Header["kid"].(string)
|
||||
switch token.Method.Alg() {
|
||||
case AlgHS256:
|
||||
if sec, ok := hsKeys[kid]; ok && kid != "" {
|
||||
return sec, nil
|
||||
}
|
||||
return hsKeys[DefaultKid], nil
|
||||
case AlgRS256:
|
||||
if rsaPub == nil {
|
||||
return nil, errors.New("服务端未配置 RS256 公钥")
|
||||
}
|
||||
return rsaPub, nil
|
||||
}
|
||||
return nil, fmt.Errorf("unsupported signing method: %v", token.Header["alg"])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 签发
|
||||
|
||||
// TokenOptions 认证中心签发参数
|
||||
type TokenOptions struct {
|
||||
Alg string // HS256 / RS256;留空时优先 RS256,未配置 RSA 则回落 HS256
|
||||
UserID int // 身份 ID(OIDC sub 用字符串形式)
|
||||
Username string
|
||||
TenantID int
|
||||
UserType string
|
||||
ClientID string
|
||||
Sid string
|
||||
Scope string
|
||||
Amr string
|
||||
Subject string
|
||||
Audience []string
|
||||
Jti string // JWT ID,用于吊销(登出/踢下线)
|
||||
TTL time.Duration // 留空默认 30 分钟
|
||||
Kid string
|
||||
}
|
||||
|
||||
// SignToken 签发 token。RSA 未配置时自动回落 HS256,保证服务可启动。
|
||||
func SignToken(opt TokenOptions) (string, error) {
|
||||
ensureKeys()
|
||||
|
||||
alg := strings.ToUpper(strings.TrimSpace(opt.Alg))
|
||||
if alg == "" {
|
||||
alg = AlgRS256
|
||||
}
|
||||
if alg == AlgRS256 && rsaPriv == nil {
|
||||
alg = AlgHS256
|
||||
}
|
||||
|
||||
ttl := opt.TTL
|
||||
if ttl <= 0 {
|
||||
ttl = 30 * time.Minute
|
||||
}
|
||||
now := time.Now()
|
||||
claims := &Claims{
|
||||
UserID: opt.UserID,
|
||||
Username: opt.Username,
|
||||
TenantId: opt.TenantID,
|
||||
UserType: opt.UserType,
|
||||
ClientID: opt.ClientID,
|
||||
Sid: opt.Sid,
|
||||
Scope: opt.Scope,
|
||||
Amr: opt.Amr,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ID: opt.Jti,
|
||||
Subject: opt.Subject,
|
||||
Audience: opt.Audience,
|
||||
ExpiresAt: jwt.NewNumericDate(now.Add(ttl)),
|
||||
IssuedAt: jwt.NewNumericDate(now),
|
||||
NotBefore: jwt.NewNumericDate(now),
|
||||
},
|
||||
}
|
||||
if issuerVal != "" {
|
||||
claims.Issuer = issuerVal
|
||||
}
|
||||
if opt.Subject == "" && opt.UserID > 0 {
|
||||
claims.Subject = fmt.Sprintf("%d", opt.UserID)
|
||||
}
|
||||
|
||||
var method jwt.SigningMethod
|
||||
switch alg {
|
||||
case AlgRS256:
|
||||
method = jwt.SigningMethodRS256
|
||||
default:
|
||||
method = jwt.SigningMethodHS256
|
||||
}
|
||||
token := jwt.NewWithClaims(method, claims)
|
||||
|
||||
kid := strings.TrimSpace(opt.Kid)
|
||||
if kid == "" && alg == AlgRS256 {
|
||||
kid = rsaKid
|
||||
}
|
||||
if kid != "" {
|
||||
token.Header["kid"] = kid
|
||||
}
|
||||
|
||||
var key interface{}
|
||||
if alg == AlgRS256 {
|
||||
key = rsaPriv
|
||||
} else {
|
||||
kidToUse := kid
|
||||
if kidToUse == "" {
|
||||
kidToUse = DefaultKid
|
||||
}
|
||||
if sec, ok := hsKeys[kidToUse]; ok {
|
||||
key = sec
|
||||
} else {
|
||||
key = hsKeys[DefaultKid]
|
||||
}
|
||||
}
|
||||
return token.SignedString(key)
|
||||
}
|
||||
|
||||
// GenerateToken 生成 JWT token(兼容旧签名与行为)。
|
||||
// 默认 HS256;若 app.conf 配置了 jwt_issuer 则带上 iss。
|
||||
func GenerateToken(userID int, username string, tenantId int, userType string) (string, error) {
|
||||
ttl := 24 * time.Hour
|
||||
ensureKeys()
|
||||
now := time.Now()
|
||||
claims := &Claims{
|
||||
UserID: userID,
|
||||
Username: username,
|
||||
TenantId: tenantId,
|
||||
UserType: userType,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(now.Add(ttl)),
|
||||
IssuedAt: jwt.NewNumericDate(now),
|
||||
NotBefore: jwt.NewNumericDate(now),
|
||||
},
|
||||
}
|
||||
if issuerVal != "" {
|
||||
claims.Issuer = issuerVal
|
||||
}
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return token.SignedString(hsKeys[DefaultKid])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------- 解析
|
||||
|
||||
// ParseToken 解析并校验 JWT。自动识别 HS256 / RS256,函数签名与旧版一致。
|
||||
func ParseToken(tokenString string) (*Claims, error) {
|
||||
claims := &Claims{}
|
||||
_, err := jwt.ParseWithClaims(tokenString, claims, keyFunc,
|
||||
jwt.WithValidMethods([]string{AlgHS256, AlgRS256}))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
// ParseTokenRaw 解析 token 但不校验有效期,用于登出/吊销场景获取 jti。
|
||||
func ParseTokenRaw(tokenString string) (*Claims, error) {
|
||||
claims := &Claims{}
|
||||
parser := jwt.NewParser(jwt.WithValidMethods([]string{AlgHS256, AlgRS256}), jwt.WithoutClaimsValidation())
|
||||
if _, err := parser.ParseWithClaims(tokenString, claims, keyFunc); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
// Issuer 返回配置的签发者(未配置时为空)
|
||||
func Issuer() string {
|
||||
ensureKeys()
|
||||
return issuerVal
|
||||
}
|
||||
|
||||
// JWK JSON Web Key(RS256 公钥)
|
||||
type JWK struct {
|
||||
Kty string `json:"kty"`
|
||||
Use string `json:"use"`
|
||||
Alg string `json:"alg"`
|
||||
Kid string `json:"kid"`
|
||||
N string `json:"n"`
|
||||
E string `json:"e"`
|
||||
}
|
||||
|
||||
// JWKS 返回 RSA 公钥集合(供 /auth/jwks.json 暴露)。
|
||||
// 各应用本地用公钥验签即可,无需每次回调认证中心的 introspect 接口。
|
||||
func JWKS() []JWK {
|
||||
ensureKeys()
|
||||
if rsaPub == nil || rsaPub.N == nil {
|
||||
return nil
|
||||
}
|
||||
return []JWK{{
|
||||
Kty: "RSA",
|
||||
Use: "sig",
|
||||
Alg: AlgRS256,
|
||||
Kid: rsaKid,
|
||||
N: base64.RawURLEncoding.EncodeToString(rsaPub.N.Bytes()),
|
||||
E: base64.RawURLEncoding.EncodeToString(big.NewInt(int64(rsaPub.E)).Bytes()),
|
||||
}}
|
||||
}
|
||||
|
||||
// HasRSA 是否已配置 RS256 密钥对(决定认证中心能否签发非对称 token)
|
||||
func HasRSA() bool {
|
||||
ensureKeys()
|
||||
return rsaPriv != nil && rsaPub != nil
|
||||
}
|
||||
|
||||
// Kid 返回当前 RS256 密钥 ID
|
||||
func Kid() string {
|
||||
ensureKeys()
|
||||
return rsaKid
|
||||
}
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
package jwtutil
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// TestGenerateParseBackwardCompat 旧签发方式产生的 token 必须仍能被解析(向后兼容)
|
||||
func TestGenerateParseBackwardCompat(t *testing.T) {
|
||||
token, err := GenerateToken(1001, "zhangsan", 7, "backend")
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken 失败: %v", err)
|
||||
}
|
||||
claims, err := ParseToken(token)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseToken 失败: %v", err)
|
||||
}
|
||||
if claims.UserID != 1001 || claims.Username != "zhangsan" || claims.TenantId != 7 || claims.UserType != "backend" {
|
||||
t.Fatalf("claims 解析不符: %+v", claims)
|
||||
}
|
||||
// 旧 token 无认证中心扩展字段
|
||||
if claims.Sid != "" || claims.ClientID != "" {
|
||||
t.Fatal("旧 token 不应携带 sid/client_id")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSignTokenFallback 未配置 RSA 时自动回落 HS256,且可解析
|
||||
func TestSignTokenFallback(t *testing.T) {
|
||||
token, err := SignToken(TokenOptions{
|
||||
Alg: AlgRS256, // 期望回落到 HS256
|
||||
UserID: 42,
|
||||
Username: "u",
|
||||
TenantID: 3,
|
||||
UserType: "tenant",
|
||||
ClientID: "crm",
|
||||
Sid: "sid-001",
|
||||
Subject: "42",
|
||||
TTL: 15 * time.Minute,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SignToken 失败: %v", err)
|
||||
}
|
||||
claims, err := ParseToken(token)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseToken 失败: %v", err)
|
||||
}
|
||||
if claims.ClientID != "crm" || claims.Sid != "sid-001" || claims.Subject != "42" {
|
||||
t.Fatalf("扩展 claims 解析不符: %+v", claims)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseInvalid 非法 token 必须报错
|
||||
func TestParseInvalid(t *testing.T) {
|
||||
if _, err := ParseToken("not-a-token"); err == nil {
|
||||
t.Fatal("非法 token 应返回错误")
|
||||
}
|
||||
if _, err := ParseToken(""); err == nil {
|
||||
t.Fatal("空 token 应返回错误")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHasRSA 无配置时不应 panic
|
||||
func TestHasRSA(t *testing.T) {
|
||||
if HasRSA() {
|
||||
t.Log("已配置 RSA 密钥对")
|
||||
}
|
||||
if Kid() == "" {
|
||||
t.Fatal("kid 不应为空")
|
||||
}
|
||||
}
|
||||
+149
-55
@@ -1,55 +1,149 @@
|
||||
package passwordutil
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
saltBytes = 16
|
||||
separator = "$"
|
||||
hashLength = 64 // sha256 hex length
|
||||
)
|
||||
|
||||
// Hash 生成 salt+hash 的存储串,格式:salt$hash(均为 hex)
|
||||
func Hash(plain string) (string, error) {
|
||||
plain = strings.TrimSpace(plain)
|
||||
if plain == "" {
|
||||
return "", errors.New("password 不能为空")
|
||||
}
|
||||
salt := make([]byte, saltBytes)
|
||||
if _, err := rand.Read(salt); err != nil {
|
||||
return "", err
|
||||
}
|
||||
saltHex := hex.EncodeToString(salt)
|
||||
hashHex := hashHex(saltHex, plain)
|
||||
return saltHex + separator + hashHex, nil
|
||||
}
|
||||
|
||||
// Verify 校验存储串(salt$hash)是否匹配输入明文密码。
|
||||
func Verify(stored, plain string) bool {
|
||||
stored = strings.TrimSpace(stored)
|
||||
plain = strings.TrimSpace(plain)
|
||||
if stored == "" || plain == "" {
|
||||
return false
|
||||
}
|
||||
parts := strings.Split(stored, separator)
|
||||
if len(parts) != 2 {
|
||||
return false
|
||||
}
|
||||
saltHex := strings.TrimSpace(parts[0])
|
||||
hashHexStored := strings.TrimSpace(parts[1])
|
||||
if saltHex == "" || len(hashHexStored) != hashLength {
|
||||
return false
|
||||
}
|
||||
return hashHex(saltHex, plain) == strings.ToLower(hashHexStored)
|
||||
}
|
||||
|
||||
func hashHex(saltHex, plain string) string {
|
||||
sum := sha256.Sum256([]byte(saltHex + plain))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
package passwordutil
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/crypto/argon2"
|
||||
)
|
||||
|
||||
// 算法标识。
|
||||
// 历史数据用 legacy(sha256 单轮,无迭代拉伸,抗 GPU 爆破能力弱);
|
||||
// 新密码统一用 argon2id,登录成功检测到 legacy 时应重新哈希升级。
|
||||
const (
|
||||
AlgoArgon2id = "argon2id"
|
||||
AlgoLegacy = "legacy"
|
||||
)
|
||||
|
||||
// argon2id 参数(OWASP 推荐起步配置:64MB / 3 轮 / 并行 2)
|
||||
const (
|
||||
argonTime = 3
|
||||
argonMemory = 64 * 1024 // 单位 KB,即 64MB
|
||||
argonThreads = 2
|
||||
argonKeyLen = 32
|
||||
argonSaltLen = 16
|
||||
)
|
||||
|
||||
// legacy 兼容参数
|
||||
const (
|
||||
saltBytes = 16
|
||||
separator = "$"
|
||||
hashLength = 64 // sha256 hex length
|
||||
)
|
||||
|
||||
// Hash 使用当前默认算法(argon2id)生成密码存储串。
|
||||
//
|
||||
// 返回 PHC 标准格式(盐与参数内联,无需独立 salt 列):
|
||||
//
|
||||
// $argon2id$v=19$m=65536,t=3,p=2$<base64(salt)>$<base64(hash)>
|
||||
//
|
||||
// 注意:函数签名与旧版一致,全部调用点无需改动即可切换到新算法。
|
||||
func Hash(plain string) (string, error) {
|
||||
plain = strings.TrimSpace(plain)
|
||||
if plain == "" {
|
||||
return "", errors.New("password 不能为空")
|
||||
}
|
||||
salt := make([]byte, argonSaltLen)
|
||||
if _, err := rand.Read(salt); err != nil {
|
||||
return "", err
|
||||
}
|
||||
key := argon2.IDKey([]byte(plain), salt, argonTime, argonMemory, argonThreads, argonKeyLen)
|
||||
return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s",
|
||||
argon2.Version,
|
||||
argonMemory, argonTime, argonThreads,
|
||||
base64.RawStdEncoding.EncodeToString(salt),
|
||||
base64.RawStdEncoding.EncodeToString(key),
|
||||
), nil
|
||||
}
|
||||
|
||||
// Verify 校验明文密码是否匹配存储串,自动识别算法:
|
||||
// - $argon2id$... → argon2id(新)
|
||||
// - <salt>$<hash> → legacy sha256(旧,兼容)
|
||||
func Verify(stored, plain string) bool {
|
||||
stored = strings.TrimSpace(stored)
|
||||
plain = strings.TrimSpace(plain)
|
||||
if stored == "" || plain == "" {
|
||||
return false
|
||||
}
|
||||
if strings.HasPrefix(stored, "$argon2id$") {
|
||||
return verifyArgon2id(stored, plain)
|
||||
}
|
||||
return verifyLegacy(stored, plain)
|
||||
}
|
||||
|
||||
// NeedsRehash 判断已存储的密码是否需要用当前算法重新哈希。
|
||||
// 登录成功后调用:返回 true 时应拿当次登录的明文重新 Hash 并落库,
|
||||
// 实现旧算法用户「首次登录自动升级」,无需强制全员重置密码。
|
||||
func NeedsRehash(stored string) bool {
|
||||
return !strings.HasPrefix(strings.TrimSpace(stored), "$argon2id$")
|
||||
}
|
||||
|
||||
// AlgoOf 返回存储串使用的算法标识,用于审计/统计旧算法存量。
|
||||
func AlgoOf(stored string) string {
|
||||
if strings.HasPrefix(strings.TrimSpace(stored), "$argon2id$") {
|
||||
return AlgoArgon2id
|
||||
}
|
||||
return AlgoLegacy
|
||||
}
|
||||
|
||||
// verifyArgon2id 解析 PHC 串并用相同参数重算比对(恒定时间比较)。
|
||||
func verifyArgon2id(stored, plain string) bool {
|
||||
parts := strings.Split(stored, "$")
|
||||
// ["", "argon2id", "v=19", "m=...,t=...,p=...", salt, hash]
|
||||
if len(parts) != 6 {
|
||||
return false
|
||||
}
|
||||
var memory, iterations, parallelism uint32
|
||||
if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &iterations, ¶llelism); err != nil {
|
||||
return false
|
||||
}
|
||||
if memory == 0 || iterations == 0 || parallelism == 0 || parallelism > 255 {
|
||||
return false
|
||||
}
|
||||
threads := uint8(parallelism)
|
||||
salt, err := base64.RawStdEncoding.DecodeString(parts[4])
|
||||
if err != nil || len(salt) == 0 {
|
||||
return false
|
||||
}
|
||||
want, err := base64.RawStdEncoding.DecodeString(parts[5])
|
||||
if err != nil || len(want) == 0 {
|
||||
return false
|
||||
}
|
||||
got := argon2.IDKey([]byte(plain), salt, iterations, memory, threads, uint32(len(want)))
|
||||
return subtle.ConstantTimeCompare(got, want) == 1
|
||||
}
|
||||
|
||||
// verifyLegacy 校验旧格式:salt$hash(均为 hex),sha256(salt+plain)。
|
||||
// 仅用于兼容历史数据,不再用于新密码。
|
||||
func verifyLegacy(stored, plain string) bool {
|
||||
stored = strings.TrimSpace(stored)
|
||||
plain = strings.TrimSpace(plain)
|
||||
if stored == "" || plain == "" {
|
||||
return false
|
||||
}
|
||||
parts := strings.Split(stored, separator)
|
||||
if len(parts) != 2 {
|
||||
return false
|
||||
}
|
||||
saltHex := strings.TrimSpace(parts[0])
|
||||
hashHexStored := strings.TrimSpace(parts[1])
|
||||
if saltHex == "" || len(hashHexStored) != hashLength {
|
||||
return false
|
||||
}
|
||||
if _, err := hex.DecodeString(saltHex); err != nil {
|
||||
return false
|
||||
}
|
||||
return hashHex(saltHex, plain) == strings.ToLower(hashHexStored)
|
||||
}
|
||||
|
||||
// hashHex 旧算法:sha256(saltHex + plain),供 verifyLegacy 使用。
|
||||
func hashHex(saltHex, plain string) string {
|
||||
sum := sha256.Sum256([]byte(saltHex + plain))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
package passwordutil
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestHashVerifyArgon2id 新算法:哈希后可校验,错误密码拒绝
|
||||
func TestHashVerifyArgon2id(t *testing.T) {
|
||||
stored, err := Hash("Passw0rd@2026")
|
||||
if err != nil {
|
||||
t.Fatalf("Hash 失败: %v", err)
|
||||
}
|
||||
if !verifyArgon2id(stored, "Passw0rd@2026") {
|
||||
t.Fatal("正确密码应校验通过")
|
||||
}
|
||||
if verifyArgon2id(stored, "wrong-password") {
|
||||
t.Fatal("错误密码不应通过")
|
||||
}
|
||||
if NeedsRehash(stored) {
|
||||
t.Fatal("argon2id 不应需要重新哈希")
|
||||
}
|
||||
if AlgoOf(stored) != AlgoArgon2id {
|
||||
t.Fatalf("算法标识应为 argon2id,实际 %s", AlgoOf(stored))
|
||||
}
|
||||
}
|
||||
|
||||
// TestLegacyCompat 旧算法兼容:历史 salt$sha256 串仍可登录,且被标记为需要升级
|
||||
func TestLegacyCompat(t *testing.T) {
|
||||
saltHex := "0123456789abcdef0123456789abcdef"
|
||||
sum := sha256.Sum256([]byte(saltHex + "oldpass"))
|
||||
legacy := saltHex + "$" + hex.EncodeToString(sum[:])
|
||||
|
||||
if !Verify(legacy, "oldpass") {
|
||||
t.Fatal("历史密码应校验通过(兼容)")
|
||||
}
|
||||
if Verify(legacy, "other") {
|
||||
t.Fatal("错误密码不应通过")
|
||||
}
|
||||
if !NeedsRehash(legacy) {
|
||||
t.Fatal("历史密码应标记为需要重新哈希")
|
||||
}
|
||||
if AlgoOf(legacy) != AlgoLegacy {
|
||||
t.Fatalf("算法标识应为 legacy,实际 %s", AlgoOf(legacy))
|
||||
}
|
||||
}
|
||||
|
||||
// TestDifferentSaltSamePassword 相同明文两次哈希结果必须不同(盐随机)
|
||||
func TestDifferentSaltSamePassword(t *testing.T) {
|
||||
a, _ := Hash("same-password")
|
||||
b, _ := Hash("same-password")
|
||||
if a == b {
|
||||
t.Fatal("相同明文两次哈希不应相同(盐必须随机)")
|
||||
}
|
||||
if !Verify(a, "same-password") || !Verify(b, "same-password") {
|
||||
t.Fatal("两份哈希都应能通过校验")
|
||||
}
|
||||
}
|
||||
|
||||
// TestEmptyPassword 空密码边界
|
||||
func TestEmptyPassword(t *testing.T) {
|
||||
if _, err := Hash(" "); err == nil {
|
||||
t.Fatal("空密码应返回错误")
|
||||
}
|
||||
if Verify("", "x") || Verify("x", "") {
|
||||
t.Fatal("空输入不应通过校验")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user