218 lines
5.7 KiB
Go
218 lines
5.7 KiB
Go
package auth
|
|
|
|
import (
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"errors"
|
|
"strconv"
|
|
"time"
|
|
|
|
"github.com/google/uuid"
|
|
|
|
"server/models"
|
|
"server/pkg/jwtutil"
|
|
)
|
|
|
|
// 认证中心错误定义
|
|
var (
|
|
ErrSessionLimitExceeded = errors.New("同时在线设备数已达上限")
|
|
ErrSessionRevoked = errors.New("会话已失效,请重新登录")
|
|
ErrSessionExpired = errors.New("会话已过期,请重新登录")
|
|
ErrRefreshTokenInvalid = errors.New("刷新令牌无效或已过期")
|
|
ErrRefreshTokenReused = errors.New("刷新令牌已被使用,疑似重放攻击,已吊销该登录")
|
|
)
|
|
|
|
// TokenPair 签发的令牌对
|
|
type TokenPair struct {
|
|
AccessToken string `json:"access_token"`
|
|
RefreshToken string `json:"refresh_token"`
|
|
TokenType string `json:"token_type"`
|
|
ExpiresIn int `json:"expires_in"`
|
|
Sid string `json:"sid"`
|
|
}
|
|
|
|
// TokenIssue 签发令牌的入参
|
|
type TokenIssue struct {
|
|
IdentityID uint64
|
|
Tid uint64
|
|
ClientID string
|
|
Sid string
|
|
Username string
|
|
UserType string
|
|
Amr string
|
|
AccessTTL int // 秒
|
|
RefreshTTL int // 秒
|
|
}
|
|
|
|
// IssueTokens 签发访问令牌与刷新令牌。
|
|
// 刷新令牌明文只在本次返回,库中仅存哈希。
|
|
func IssueTokens(opt TokenIssue) (*TokenPair, error) {
|
|
accessTTL := opt.AccessTTL
|
|
if accessTTL <= 0 {
|
|
accessTTL = 1800
|
|
}
|
|
refreshTTL := opt.RefreshTTL
|
|
if refreshTTL <= 0 {
|
|
refreshTTL = 2592000
|
|
}
|
|
|
|
jti := uuid.NewString()
|
|
access, err := jwtutil.SignToken(jwtutil.TokenOptions{
|
|
Alg: jwtutil.AlgRS256,
|
|
UserID: int(opt.IdentityID),
|
|
Username: opt.Username,
|
|
TenantID: int(opt.Tid),
|
|
UserType: opt.UserType,
|
|
ClientID: opt.ClientID,
|
|
Sid: opt.Sid,
|
|
Amr: opt.Amr,
|
|
Subject: strconv.FormatUint(opt.IdentityID, 10),
|
|
Audience: []string{opt.ClientID},
|
|
Jti: jti,
|
|
TTL: time.Duration(accessTTL) * time.Second,
|
|
})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
plain, err := randomToken(32)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
family := uuid.NewString()
|
|
now := time.Now()
|
|
rt := &models.AuthRefreshToken{
|
|
TokenHash: hashToken(plain),
|
|
IdentityID: opt.IdentityID,
|
|
Tid: opt.Tid,
|
|
ClientID: opt.ClientID,
|
|
Sid: opt.Sid,
|
|
FamilyID: family,
|
|
ExpiresAt: now.Add(time.Duration(refreshTTL) * time.Second),
|
|
}
|
|
if _, err := models.Orm.Insert(rt); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return &TokenPair{
|
|
AccessToken: access,
|
|
RefreshToken: plain,
|
|
TokenType: "Bearer",
|
|
ExpiresIn: accessTTL,
|
|
Sid: opt.Sid,
|
|
}, nil
|
|
}
|
|
|
|
// RefreshTokens 用刷新令牌换取新的令牌对(轮换 + 重放检测)。
|
|
//
|
|
// 安全要点:检测到已使用的刷新令牌再次出现时,判定为重放,
|
|
// 吊销整个 family(该次登录的全部令牌)。
|
|
func RefreshTokens(plain, clientID string) (*TokenPair, error) {
|
|
h := hashToken(plain)
|
|
var rt models.AuthRefreshToken
|
|
if err := models.Orm.QueryTable(new(models.AuthRefreshToken)).Filter("token_hash", h).One(&rt); err != nil {
|
|
return nil, ErrRefreshTokenInvalid
|
|
}
|
|
if rt.ClientID != "" && clientID != "" && rt.ClientID != clientID {
|
|
return nil, ErrRefreshTokenInvalid
|
|
}
|
|
if rt.Revoked != 0 || rt.ExpiresAt.Before(time.Now()) {
|
|
return nil, ErrRefreshTokenInvalid
|
|
}
|
|
if rt.Used != 0 {
|
|
// 重放:吊销同族全部令牌与该会话
|
|
_, _ = models.Orm.QueryTable(new(models.AuthRefreshToken)).
|
|
Filter("family_id", rt.FamilyID).
|
|
Update(map[string]interface{}{"revoked": 1})
|
|
_ = RevokeSession(rt.Sid, models.RevokeReasonAdmin)
|
|
return nil, ErrRefreshTokenReused
|
|
}
|
|
|
|
// 标记已用(保留 Row 以便审计)
|
|
_, _ = models.Orm.QueryTable(new(models.AuthRefreshToken)).
|
|
Filter("id", rt.ID).
|
|
Update(map[string]interface{}{"used": 1})
|
|
|
|
// 会话校验:已吊销/过期则拒绝续期
|
|
s, err := GetSession(rt.Sid)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
_ = TouchSession(rt.Sid)
|
|
|
|
pair, err := IssueTokens(TokenIssue{
|
|
IdentityID: rt.IdentityID,
|
|
Tid: s.Tid, // 以会话当前企业为准(支持切换企业后刷新)
|
|
ClientID: rt.ClientID,
|
|
Sid: rt.Sid,
|
|
Amr: derefStr(s.Amr),
|
|
})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
// 新令牌继承同一 family,便于后续溯源与整族吊销
|
|
if _, err := models.Orm.QueryTable(new(models.AuthRefreshToken)).
|
|
Filter("token_hash", hashToken(pair.RefreshToken)).
|
|
Update(map[string]interface{}{"family_id": rt.FamilyID, "rotated_from": h}); err != nil {
|
|
return nil, err
|
|
}
|
|
return pair, nil
|
|
}
|
|
|
|
// RevokeTokenPair 登出:吊销刷新令牌、会话,并把 access token 的 jti 加入黑名单。
|
|
func RevokeTokenPair(plain, accessToken, reason string) error {
|
|
if plain != "" {
|
|
_, _ = models.Orm.QueryTable(new(models.AuthRefreshToken)).
|
|
Filter("token_hash", hashToken(plain)).
|
|
Update(map[string]interface{}{"revoked": 1})
|
|
}
|
|
if accessToken != "" {
|
|
if claims, err := jwtutil.ParseTokenRaw(accessToken); err == nil {
|
|
_ = Blacklist(claims.ID, claims.Sid, reason, time.Unix(claims.ExpiresAt.Unix(), 0))
|
|
if claims.Sid != "" {
|
|
_ = RevokeSession(claims.Sid, reason)
|
|
}
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// Blacklist 把 jti 加入吊销表
|
|
func Blacklist(jti, sid, reason string, expiresAt time.Time) error {
|
|
if jti == "" {
|
|
return nil
|
|
}
|
|
item := &models.AuthTokenBlacklist{
|
|
Jti: jti,
|
|
ExpiresAt: expiresAt,
|
|
}
|
|
if sid != "" {
|
|
item.Sid = &sid
|
|
}
|
|
if reason != "" {
|
|
item.Reason = &reason
|
|
}
|
|
_, err := models.Orm.InsertOrUpdate(item, "jti")
|
|
return err
|
|
}
|
|
|
|
// IsBlacklisted 判断 jti 是否已被吊销
|
|
func IsBlacklisted(jti string) bool {
|
|
if jti == "" {
|
|
return false
|
|
}
|
|
return models.Orm.QueryTable(new(models.AuthTokenBlacklist)).Filter("jti", jti).Exist()
|
|
}
|
|
|
|
func hashToken(plain string) string {
|
|
sum := sha256.Sum256([]byte(plain))
|
|
return hex.EncodeToString(sum[:])
|
|
}
|
|
|
|
func derefStr(p *string) string {
|
|
if p == nil {
|
|
return ""
|
|
}
|
|
return *p
|
|
}
|