Files
yunzerwebsiteallinone/go/services/auth/token.go
T
2026-09-19 21:44:04 +08:00

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
}