434 lines
14 KiB
Go
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)
|
|
}
|
|
}
|