Files
2026-07-30 12:23:05 +08:00

434 lines
14 KiB
Go

package service
import (
"context"
"errors"
"fmt"
"log"
"sync"
"time"
"github.com/golang-jwt/jwt/v5"
"gorm.io/gorm"
"oci-portal/internal/model"
)
// ErrSessionCurrent 表示试图撤销当前会话;应引导用退出登录。
var ErrSessionCurrent = errors.New("不能撤销当前会话,请使用退出登录")
const (
// sessionSeenTTL 是最后活跃时间的回写节流窗口;窗口内命中缓存直接放行。
sessionSeenTTL = time.Minute
// sessionCleanupTick 是会话行清理周期。
sessionCleanupTick = time.Hour
// sessionExpiredKeep / sessionRevokedKeep 是失效行的保留期,过后物理删除。
sessionExpiredKeep = 24 * time.Hour
sessionRevokedKeep = 7 * 24 * time.Hour
)
// jtiTombstones 是定点撤销的负缓存:无容量上限,不存在「淘汰导致已撤销
// 令牌复活」;增长由 TTL 清理约束——撤销是认证后的低频人工操作,
// 集合尺寸恒小,put 时线性清理过期项即可。
type jtiTombstones struct {
mu sync.RWMutex
m map[string]time.Time // jti → 过期时刻
}
func newJtiTombstones() *jtiTombstones {
return &jtiTombstones{m: map[string]time.Time{}}
}
// put 登记撤销标记并顺带清理过期项。
func (t *jtiTombstones) put(jti string, ttl time.Duration) {
if jti == "" || ttl <= 0 {
return
}
now := time.Now()
t.mu.Lock()
defer t.mu.Unlock()
for k, exp := range t.m {
if now.After(exp) {
delete(t.m, k)
}
}
t.m[jti] = now.Add(ttl)
}
// has 报告 jti 是否在有效撤销标记中。
func (t *jtiTombstones) has(jti string) bool {
t.mu.RLock()
exp, ok := t.m[jti]
t.mu.RUnlock()
return ok && time.Now().Before(exp)
}
// SessionMeta 是签发会话时的客户端上下文;Method 由各登录出口的 service 层填写,
// api 层只采集 IP/UA。零值 meta 表示不落会话行(内部/测试场景)。
type SessionMeta struct {
ClientIP string
UserAgent string
Method string // password / oidc / github / passkey / wallet
}
// empty 报告 meta 是否为「不落行」哨兵。
func (m SessionMeta) empty() bool { return m.ClientIP == "" && m.UserAgent == "" }
// signSessionToken 签发 JWT 并按 meta 落地会话行。
func (s *AuthService) signSessionToken(ctx context.Context, user *model.User, meta SessionMeta) (string, time.Time, error) {
return s.signSessionTokenDB(s.db.WithContext(ctx), user, meta)
}
// signSessionTokenTx 是身份登录事务内的签发入口。
func (s *AuthService) signSessionTokenTx(tx *gorm.DB, user *model.User, meta SessionMeta) (string, time.Time, error) {
return s.signSessionTokenDB(tx, user, meta)
}
func (s *AuthService) signSessionTokenDB(db *gorm.DB, user *model.User, meta SessionMeta) (string, time.Time, error) {
token, expires, jti, err := s.signToken(user.Username, user.TokenVersion)
if err != nil || meta.empty() {
return token, expires, err
}
now := time.Now()
row := model.UserSession{
UserID: user.ID, TokenID: jti, TokenVer: user.TokenVersion,
Method: meta.Method, ClientIP: meta.ClientIP, UserAgent: meta.UserAgent,
LastSeenAt: now, ExpiresAt: expires,
}
if err := db.Create(&row).Error; err != nil {
// fail-closed:落行失败拒发令牌,否则产生列表不可见、无法定点撤销的孤儿会话
return "", time.Time{}, fmt.Errorf("record session: %w", err)
}
return token, expires, nil
}
// RenewToken 为敏感变更后的操作者换发新令牌:旧令牌对应的会话行接续
// (沿用 jti,同行更新版本/有效期,保留登录方式与创建时间),无行则按 meta 新建。
// 旧令牌只验签名不验有效性——版本刚被递增,旧令牌语义上已失效但行仍需接续。
func (s *AuthService) RenewToken(ctx context.Context, username, oldToken string, meta SessionMeta) (string, time.Time, error) {
proof, ok := s.signedTokenProof(oldToken)
if !ok {
return "", time.Time{}, ErrTokenStale
}
var token string
var expires time.Time
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
user, err := lockUserForAuthChange(tx, username)
if err != nil {
return err
}
if user.TokenVersion != proof.Ver+1 {
return ErrTokenStale
}
token, expires, err = s.renewSessionTx(tx, user, oldToken, meta)
return err
})
return token, expires, err
}
// renewSessionTx 事务内签新令牌并接续会话行(RenewToken 的事务内核):
// 与认证因子写入 / 版本递增同事务提交,绑定等敏感变更全程要么全成要么全滚;
// 调用方须保证 user.TokenVersion 已是递增后的最新值。
func (s *AuthService) renewSessionTx(tx *gorm.DB, user *model.User, oldToken string, meta SessionMeta) (string, time.Time, error) {
jti := s.signedJti(oldToken)
if jti == "" {
jti = newTokenID()
}
token, expires, _, err := s.signTokenWithJTI(user.Username, user.TokenVersion, jti)
if err != nil {
return "", time.Time{}, err
}
updates := map[string]any{
"token_id": jti, "token_ver": user.TokenVersion,
"expires_at": expires, "last_seen_at": time.Now(),
}
renewed, err := s.renewExistingSessionTx(tx, user.ID, oldToken, updates)
if err != nil {
return "", time.Time{}, err
}
if renewed {
return token, expires, nil
}
if meta.empty() {
return token, expires, nil
}
err = createRenewedSessionTx(tx, user, jti, expires, meta)
return token, expires, err
}
func createRenewedSessionTx(
tx *gorm.DB, user *model.User, jti string, expires time.Time, meta SessionMeta,
) error {
row := model.UserSession{
UserID: user.ID, TokenID: jti, TokenVer: user.TokenVersion,
Method: meta.Method, ClientIP: meta.ClientIP, UserAgent: meta.UserAgent,
LastSeenAt: time.Now(), ExpiresAt: expires,
}
if err := tx.Create(&row).Error; err != nil {
return fmt.Errorf("record renewed session: %w", err)
}
return nil
}
// renewExistingSessionTx 仅在旧行仍有效时接续;有行但已撤销或更新失败均失败关闭。
func (s *AuthService) renewExistingSessionTx(tx *gorm.DB, userID uint, oldToken string, updates map[string]any) (bool, error) {
oldJTI := s.signedJti(oldToken)
if oldJTI == "" {
return false, nil
}
if s.revokedJti.has(oldJTI) {
return false, ErrTokenStale
}
if _, hit := s.revoked.Get(tokenHash(oldToken)); hit {
return false, ErrTokenStale
}
res := tx.Model(&model.UserSession{}).
Where("token_id = ? AND user_id = ? AND revoked_at IS NULL", oldJTI, userID).Updates(updates)
if res.Error != nil {
return false, fmt.Errorf("renew session: %w", res.Error)
}
if res.RowsAffected > 0 {
s.seen.DeletePrefix("seen|" + oldJTI)
return true, nil
}
return false, s.ensureRenewalHasNoOldRow(tx, userID, oldJTI)
}
func (s *AuthService) ensureRenewalHasNoOldRow(tx *gorm.DB, userID uint, oldJTI string) error {
var count int64
err := tx.Model(&model.UserSession{}).
Where("token_id = ? AND user_id = ?", oldJTI, userID).Count(&count).Error
if err != nil {
return fmt.Errorf("check renewed session: %w", err)
}
if count > 0 {
return ErrTokenStale
}
return nil
}
// signedJti 校验令牌签名并取出 jti;不验有效期与版本(换发场景旧令牌刚失效)。
// 签名必须有效,防止伪造 jti 抢占他人会话行。
func (s *AuthService) signedJti(tokenString string) string {
proof, ok := s.signedTokenProof(tokenString)
if !ok {
return ""
}
return proof.Jti
}
func (s *AuthService) signedTokenProof(tokenString string) (TokenProof, bool) {
if tokenString == "" {
return TokenProof{}, false
}
claims := &authClaims{}
_, err := jwt.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (any, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, fmt.Errorf("unexpected signing method %v", t.Header["alg"])
}
return s.jwtSecret, nil
}, jwt.WithoutClaimsValidation())
if err != nil {
return TokenProof{}, false
}
return TokenProof{Ver: claims.Ver, Jti: claims.ID}, true
}
// checkSession 校验 jti 对应会话未被定点撤销;无行放行(存量令牌兼容)。
// 有效会话按节流窗口回写最后活跃时间;撤销动作会清掉节流缓存保证即时生效。
func (s *AuthService) checkSession(ctx context.Context, jti string) error {
if jti == "" {
return nil
}
// 撤销负缓存优先:防「读库通过→撤销→回写 seen」竞态让已撤销令牌复活
if s.revokedJti.has(jti) {
return errors.New("session revoked")
}
if _, hit := s.seen.Get("seen|" + jti); hit {
return nil
}
var row model.UserSession
err := s.db.WithContext(ctx).Select("id", "revoked_at").Where("token_id = ?", jti).First(&row).Error
if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
return fmt.Errorf("find session: %w", err)
}
if row.RevokedAt != nil {
return errors.New("session revoked")
}
s.seen.Set("seen|"+jti, struct{}{}, sessionSeenTTL)
s.db.WithContext(ctx).Model(&model.UserSession{}).
Where("id = ?", row.ID).UpdateColumn("last_seen_at", time.Now())
return nil
}
// SessionInfo 是会话列表条目;Current 标记请求者自身会话。
type SessionInfo struct {
model.UserSession
Current bool `json:"current"`
}
// ListSessions 列出账号的活跃会话(未撤销、未过期、版本为当前),最近活跃在前。
func (s *AuthService) ListSessions(ctx context.Context, username, currentToken string) ([]SessionInfo, error) {
user, err := s.findUser(ctx, username)
if err != nil {
return nil, err
}
rows := []model.UserSession{}
err = s.db.WithContext(ctx).
Where("user_id = ? AND revoked_at IS NULL AND expires_at > ? AND token_ver = ?",
user.ID, time.Now(), user.TokenVersion).
Order("last_seen_at DESC").Find(&rows).Error
if err != nil {
return nil, fmt.Errorf("list sessions: %w", err)
}
currentJti := s.signedJti(currentToken)
out := make([]SessionInfo, 0, len(rows))
for _, r := range rows {
out = append(out, SessionInfo{UserSession: r, Current: r.TokenID == currentJti})
}
return out, nil
}
// RevokeSession 定点撤销一个会话(校验归属);当前会话拒绝(引导登出),
// 不递增令牌版本、不影响其余会话。
func (s *AuthService) RevokeSession(ctx context.Context, username, currentToken string, id uint) error {
currentJTI := s.signedJti(currentToken)
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
user, err := lockUserForAuthChange(tx, username)
if err != nil {
return err
}
row, err := sessionForRevokeTx(tx, user.ID, id)
if err != nil {
return err
}
if row.TokenID == currentJTI {
return ErrSessionCurrent
}
return s.revokeSessionRowTx(tx, row)
})
}
func sessionForRevokeTx(tx *gorm.DB, userID, id uint) (*model.UserSession, error) {
var row model.UserSession
err := tx.Where("id = ? AND user_id = ? AND revoked_at IS NULL", id, userID).First(&row).Error
return &row, err
}
func (s *AuthService) revokeSessionRowTx(tx *gorm.DB, row *model.UserSession) error {
res := tx.Model(&model.UserSession{}).
Where("id = ? AND revoked_at IS NULL", row.ID).UpdateColumn("revoked_at", time.Now())
if res.Error != nil {
return fmt.Errorf("revoke session: %w", res.Error)
}
if res.RowsAffected == 0 {
return gorm.ErrRecordNotFound
}
s.markJTIRevoked(row.TokenID, time.Until(row.ExpiresAt))
return nil
}
// revokeSessionByJTI 按 jti 标记会话撤销(登出联动);无行为无害操作。
func (s *AuthService) revokeSessionByJTI(ctx context.Context, username, jti string, ttl time.Duration) {
s.markJTIRevoked(jti, logoutJTITTL(ttl))
if jti == "" {
return
}
userID := s.logoutUserID(ctx, username, jti)
if userID == 0 {
return
}
_ = s.revokeJTIForUser(ctx, userID, jti)
}
// logoutJTITTL 覆盖旧令牌剩余窗口内最晚产生的同 JTI 换发令牌。
func logoutJTITTL(ttl time.Duration) time.Duration { return ttl + tokenTTL }
func (s *AuthService) logoutUserID(ctx context.Context, username, jti string) uint {
if userID := s.sessionUserID(ctx, jti); userID != 0 {
return userID
}
if userID := s.usernameUserID(ctx, username); userID != 0 {
return userID
}
// legacy 换发行可能正提交;再读一次缩小「无行→建行」窗口。
return s.sessionUserID(ctx, jti)
}
func (s *AuthService) sessionUserID(ctx context.Context, jti string) uint {
var row model.UserSession
err := s.db.WithContext(ctx).Select("user_id").
Where("token_id = ?", jti).First(&row).Error
if err != nil {
return 0
}
return row.UserID
}
func (s *AuthService) usernameUserID(ctx context.Context, username string) uint {
if username == "" {
return 0
}
var user model.User
if err := s.db.WithContext(ctx).Select("id").
Where("username = ?", username).First(&user).Error; err != nil {
return 0
}
return user.ID
}
func (s *AuthService) revokeJTIForUser(ctx context.Context, userID uint, jti string) error {
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if _, err := lockUserByIDForAuthChange(tx, userID); err != nil {
return err
}
return tx.Model(&model.UserSession{}).
Where("token_id = ? AND user_id = ? AND revoked_at IS NULL", jti, userID).
UpdateColumn("revoked_at", time.Now()).Error
})
}
func (s *AuthService) markJTIRevoked(jti string, ttl time.Duration) {
s.seen.DeletePrefix("seen|" + jti)
s.revokedJti.put(jti, ttl)
}
// StartSessionCleanup 启动会话行周期清理:启动即清一次,之后每小时一次,
// 随 ctx 取消退出;WaitSessionCleanup 可等待其真正结束(并发规范)。
func (s *AuthService) StartSessionCleanup(ctx context.Context) {
s.cleanupWG.Add(1)
go func() {
defer s.cleanupWG.Done()
s.cleanupSessionsOnce(ctx)
ticker := time.NewTicker(sessionCleanupTick)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
s.cleanupSessionsOnce(ctx)
}
}
}()
}
// WaitSessionCleanup 阻塞等待清理 goroutine 退出(取消 ctx 后调用)。
func (s *AuthService) WaitSessionCleanup() {
s.cleanupWG.Wait()
}
// cleanupSessionsOnce 删除保留期外的失效行;失败只记日志、不中断周期调度。
func (s *AuthService) cleanupSessionsOnce(ctx context.Context) {
now := time.Now()
err := s.db.WithContext(ctx).
Where("expires_at < ? OR revoked_at < ?", now.Add(-sessionExpiredKeep), now.Add(-sessionRevokedKeep)).
Delete(&model.UserSession{}).Error
if err != nil {
log.Printf("session cleanup: %v", err)
}
}