* refactor(auth): replace dashboard sessions with stateless tokens * feat(auth): harden session issuance and distributed enforcement * fix(proxy): preserve trusted proxy compatibility defaults * refactor: address dashboard auth review feedback * refactor: remove classic frontend and flatten web app
871 lines
29 KiB
Go
871 lines
29 KiB
Go
package model
|
|
|
|
import (
|
|
"context"
|
|
"crypto/hmac"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
"time"
|
|
|
|
"github.com/QuantumNous/new-api/common"
|
|
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
const (
|
|
UserSessionStatusActive = "active"
|
|
UserSessionStatusRevoking = "revoking"
|
|
UserSessionStatusRevoked = "revoked"
|
|
|
|
userSessionCacheSchema = 1
|
|
userSessionListLimit = 100
|
|
userSessionRevokeBatchSize = 500
|
|
userSessionCleanupScanLimit = 1000
|
|
userSessionCleanupBatchSize = 500
|
|
)
|
|
|
|
var (
|
|
ErrUserSessionInvalid = errors.New("user session is invalid")
|
|
ErrUserSessionInactive = errors.New("user session is inactive")
|
|
ErrUserSessionRefreshInvalid = errors.New("user session refresh token is invalid")
|
|
ErrUserSessionRefreshRace = errors.New("user session refresh is already in progress")
|
|
ErrUserSessionRefreshReuse = errors.New("user session refresh token was reused")
|
|
ErrUserSessionLimit = errors.New("active user session limit reached")
|
|
ErrUserSessionIssuanceLimit = errors.New("user session issuance limit reached")
|
|
errUserSessionCacheObservationStale = errors.New("user session cache observation is stale")
|
|
)
|
|
|
|
// UserSession is the server-side control plane for short-lived access JWTs.
|
|
// RefreshHash values are HMAC digests supplied by the service layer; opaque
|
|
// refresh secrets are never persisted.
|
|
type UserSession struct {
|
|
SID string `json:"sid" gorm:"column:sid;type:varchar(64);primaryKey"`
|
|
UserID int `json:"user_id" gorm:"column:user_id;not null;index:idx_user_sessions_user_status_expiry,priority:1;index:idx_user_sessions_user_created,priority:1"`
|
|
Version int64 `json:"version" gorm:"type:bigint;not null;default:1"`
|
|
UserAuthVersion int64 `json:"user_auth_version" gorm:"type:bigint;not null"`
|
|
Status string `json:"status" gorm:"type:varchar(16);not null;index:idx_user_sessions_user_status_expiry,priority:2;index:idx_user_sessions_status_revoked,priority:1"`
|
|
RefreshHash string `json:"-" gorm:"type:char(64);not null"`
|
|
PreviousRefreshHash string `json:"-" gorm:"type:varchar(64)"`
|
|
PreviousValidUntil int64 `json:"-" gorm:"type:bigint;not null;default:0"`
|
|
LoginMethod string `json:"login_method" gorm:"type:varchar(32);not null"`
|
|
IP string `json:"ip" gorm:"type:varchar(64)"`
|
|
UserAgent string `json:"user_agent" gorm:"type:text"`
|
|
CreatedAt int64 `json:"created_at" gorm:"autoCreateTime;column:created_at;index:idx_user_sessions_user_created,priority:2"`
|
|
LastActiveAt int64 `json:"last_active_at" gorm:"type:bigint;not null;column:last_active_at"`
|
|
ExpiresAt int64 `json:"expires_at" gorm:"type:bigint;not null;column:expires_at;index:idx_user_sessions_user_status_expiry,priority:3;index:idx_user_sessions_expires_at"`
|
|
RevokedAt int64 `json:"revoked_at,omitempty" gorm:"type:bigint;not null;default:0;column:revoked_at;index:idx_user_sessions_status_revoked,priority:2"`
|
|
RevokedReason string `json:"revoked_reason,omitempty" gorm:"type:varchar(64);column:revoked_reason"`
|
|
}
|
|
|
|
func (UserSession) TableName() string {
|
|
return "user_sessions"
|
|
}
|
|
|
|
func (session *UserSession) AfterFind(_ *gorm.DB) error {
|
|
session.PreviousRefreshHash = strings.TrimSpace(session.PreviousRefreshHash)
|
|
return nil
|
|
}
|
|
|
|
type userSessionCacheEntry struct {
|
|
SID string
|
|
UserID int
|
|
Version int64
|
|
UserAuthVersion int64
|
|
Status string
|
|
LoginMethod string
|
|
IP string
|
|
UserAgent string
|
|
CreatedAt int64
|
|
LastActiveAt int64
|
|
ExpiresAt int64
|
|
RevokedAt int64
|
|
RevokedReason string
|
|
CacheSchema int
|
|
}
|
|
|
|
func (session *UserSession) cacheEntry() *userSessionCacheEntry {
|
|
return &userSessionCacheEntry{
|
|
SID: session.SID,
|
|
UserID: session.UserID,
|
|
Version: session.Version,
|
|
UserAuthVersion: session.UserAuthVersion,
|
|
Status: session.Status,
|
|
LoginMethod: session.LoginMethod,
|
|
IP: session.IP,
|
|
UserAgent: session.UserAgent,
|
|
CreatedAt: session.CreatedAt,
|
|
LastActiveAt: session.LastActiveAt,
|
|
ExpiresAt: session.ExpiresAt,
|
|
RevokedAt: session.RevokedAt,
|
|
RevokedReason: session.RevokedReason,
|
|
CacheSchema: userSessionCacheSchema,
|
|
}
|
|
}
|
|
|
|
func (entry *userSessionCacheEntry) session() *UserSession {
|
|
return &UserSession{
|
|
SID: entry.SID,
|
|
UserID: entry.UserID,
|
|
Version: entry.Version,
|
|
UserAuthVersion: entry.UserAuthVersion,
|
|
Status: entry.Status,
|
|
LoginMethod: entry.LoginMethod,
|
|
IP: entry.IP,
|
|
UserAgent: entry.UserAgent,
|
|
CreatedAt: entry.CreatedAt,
|
|
LastActiveAt: entry.LastActiveAt,
|
|
ExpiresAt: entry.ExpiresAt,
|
|
RevokedAt: entry.RevokedAt,
|
|
RevokedReason: entry.RevokedReason,
|
|
}
|
|
}
|
|
|
|
func userSessionCacheKey(sid string) string {
|
|
digest := common.GenerateHMACWithKey([]byte("user-session-cache-v1:"+common.SessionSecret), sid)
|
|
return "auth:session:" + digest
|
|
}
|
|
|
|
func userSessionCacheDeadline() time.Time {
|
|
return time.Now().Add(time.Duration(userCacheTTLSeconds()) * time.Second)
|
|
}
|
|
|
|
func CreateUserSession(session *UserSession) error {
|
|
now := time.Now().Unix()
|
|
if session == nil || session.SID == "" || session.UserID <= 0 || session.UserAuthVersion <= 0 || session.RefreshHash == "" || session.ExpiresAt <= now {
|
|
return ErrUserSessionInvalid
|
|
}
|
|
if session.Version <= 0 {
|
|
session.Version = 1
|
|
}
|
|
if session.Status == "" {
|
|
session.Status = UserSessionStatusActive
|
|
}
|
|
if session.Status != UserSessionStatusActive || session.RevokedAt != 0 {
|
|
return ErrUserSessionInvalid
|
|
}
|
|
if session.LastActiveAt == 0 {
|
|
session.LastActiveAt = now
|
|
}
|
|
if session.CreatedAt == 0 {
|
|
session.CreatedAt = now
|
|
}
|
|
cacheDeadline := userSessionCacheDeadline()
|
|
if err := DB.Create(session).Error; err != nil {
|
|
return err
|
|
}
|
|
if err := writeUserSessionCache(session.cacheEntry(), cacheDeadline); err != nil {
|
|
if errors.Is(err, errUserSessionCacheObservationStale) {
|
|
return confirmUserSessionActiveSnapshot(session)
|
|
}
|
|
if errors.Is(err, ErrUserSessionInactive) {
|
|
return err
|
|
}
|
|
common.SysLog("failed to populate newly created user session cache: " + err.Error())
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func CountActiveUserSessions(userID int, now int64) (int64, error) {
|
|
if userID <= 0 {
|
|
return 0, ErrUserSessionInvalid
|
|
}
|
|
if now <= 0 {
|
|
now = time.Now().Unix()
|
|
}
|
|
var count int64
|
|
err := DB.Model(&UserSession{}).
|
|
Where("user_id = ? AND status = ? AND expires_at > ?", userID, UserSessionStatusActive, now).
|
|
Count(&count).Error
|
|
return count, err
|
|
}
|
|
|
|
// CountUserSessionsCreatedSince counts every issued row, regardless of its
|
|
// current status or expiry. userID zero selects the global count.
|
|
func CountUserSessionsCreatedSince(userID int, createdAfter int64) (int64, error) {
|
|
if userID < 0 || createdAfter <= 0 {
|
|
return 0, ErrUserSessionInvalid
|
|
}
|
|
query := DB.Model(&UserSession{}).Where("created_at > ?", createdAfter)
|
|
if userID > 0 {
|
|
query = query.Where("user_id = ?", userID)
|
|
}
|
|
var count int64
|
|
err := query.Count(&count).Error
|
|
return count, err
|
|
}
|
|
|
|
func GetUserSessionBySID(sid string) (*UserSession, error) {
|
|
if sid == "" {
|
|
return nil, ErrUserSessionInvalid
|
|
}
|
|
var session UserSession
|
|
if err := DB.Where("sid = ?", sid).First(&session).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
return &session, nil
|
|
}
|
|
|
|
// GetUserSessionCached validates cached state first and falls back to the
|
|
// database on a miss or Redis read failure. A deny tombstone never falls back.
|
|
func GetUserSessionCached(sid string) (*UserSession, error) {
|
|
if sid == "" {
|
|
return nil, ErrUserSessionInvalid
|
|
}
|
|
if common.RedisEnabled {
|
|
entry, err := getUserSessionCache(sid)
|
|
if err == nil {
|
|
return entry.session(), nil
|
|
}
|
|
if errors.Is(err, ErrUserSessionInactive) {
|
|
return nil, err
|
|
}
|
|
}
|
|
|
|
cacheDeadline := userSessionCacheDeadline()
|
|
session, err := GetUserSessionBySID(sid)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
now := time.Now().Unix()
|
|
if session.Status != UserSessionStatusActive || session.RevokedAt != 0 || session.ExpiresAt <= now {
|
|
if common.RedisEnabled {
|
|
entry := session.cacheEntry()
|
|
entry.Status = UserSessionStatusRevoked
|
|
_ = writeUserSessionCache(entry, time.Time{})
|
|
}
|
|
return nil, ErrUserSessionInactive
|
|
}
|
|
if common.RedisEnabled {
|
|
if err := writeUserSessionCache(session.cacheEntry(), cacheDeadline); err != nil {
|
|
if errors.Is(err, errUserSessionCacheObservationStale) {
|
|
if confirmErr := confirmUserSessionActiveSnapshot(session); confirmErr != nil {
|
|
return nil, confirmErr
|
|
}
|
|
return session, nil
|
|
}
|
|
if errors.Is(err, ErrUserSessionInactive) {
|
|
return nil, err
|
|
}
|
|
common.SysLog("failed to synchronously populate user session cache: " + err.Error())
|
|
}
|
|
}
|
|
return session, nil
|
|
}
|
|
|
|
func getUserSessionCache(sid string) (*userSessionCacheEntry, error) {
|
|
var entry userSessionCacheEntry
|
|
if err := common.RedisHGetObj(userSessionCacheKey(sid), &entry); err != nil {
|
|
return nil, err
|
|
}
|
|
if entry.CacheSchema != userSessionCacheSchema || entry.SID != sid || entry.UserID <= 0 || entry.Version <= 0 || entry.UserAuthVersion <= 0 {
|
|
return nil, fmt.Errorf("user session cache schema is stale")
|
|
}
|
|
if entry.Status != UserSessionStatusActive || entry.RevokedAt != 0 || entry.ExpiresAt <= time.Now().Unix() {
|
|
return nil, ErrUserSessionInactive
|
|
}
|
|
return &entry, nil
|
|
}
|
|
|
|
// writeUserSessionCache writes a bounded Session snapshot. Active snapshots
|
|
// must carry a deadline captured immediately before their authoritative
|
|
// database read or mutation. Delayed fills inherit the unspent portion of that
|
|
// window, so a stale active snapshot cannot outlive a short deny tombstone and
|
|
// reactivate a revoked Session after the tombstone expires. Deny states pass a
|
|
// zero deadline because their TTL starts when they are published.
|
|
func writeUserSessionCache(entry *userSessionCacheEntry, cacheDeadline time.Time) error {
|
|
if entry == nil || !common.RedisEnabled {
|
|
return nil
|
|
}
|
|
now := time.Now()
|
|
sessionExpiresAt := time.Unix(entry.ExpiresAt, 0)
|
|
sessionTTL := sessionExpiresAt.Sub(now)
|
|
var redisExpiration int64
|
|
if entry.Status == UserSessionStatusActive {
|
|
if cacheDeadline.IsZero() {
|
|
return ErrUserSessionInvalid
|
|
}
|
|
cacheTTL := cacheDeadline.Sub(now)
|
|
if cacheTTL <= 0 {
|
|
return errUserSessionCacheObservationStale
|
|
}
|
|
if sessionTTL <= 0 {
|
|
return ErrUserSessionInactive
|
|
}
|
|
cacheExpiresAt := cacheDeadline
|
|
if sessionExpiresAt.Before(cacheExpiresAt) {
|
|
cacheExpiresAt = sessionExpiresAt
|
|
}
|
|
if cacheExpiresAt.Sub(now) < time.Millisecond {
|
|
return errUserSessionCacheObservationStale
|
|
}
|
|
redisExpiration = cacheExpiresAt.UnixMilli()
|
|
} else {
|
|
ttl := min(sessionTTL, time.Duration(userCacheTTLSeconds())*time.Second)
|
|
if ttl <= 0 {
|
|
ttl = time.Second
|
|
}
|
|
redisExpiration = ttl.Milliseconds()
|
|
if redisExpiration <= 0 {
|
|
redisExpiration = 1
|
|
}
|
|
}
|
|
entry.CacheSchema = userSessionCacheSchema
|
|
const script = `
|
|
local current_status = redis.call('HGET', KEYS[1], 'Status')
|
|
local current_version = tonumber(redis.call('HGET', KEYS[1], 'Version') or '0')
|
|
if ARGV[5] == 'active' and (current_status == 'revoking' or current_status == 'revoked') then
|
|
return 0
|
|
end
|
|
if current_version > tonumber(ARGV[3]) then
|
|
return 0
|
|
end
|
|
redis.call('HSET', KEYS[1],
|
|
'SID', ARGV[1], 'UserID', ARGV[2], 'Version', ARGV[3],
|
|
'UserAuthVersion', ARGV[4], 'Status', ARGV[5],
|
|
'LoginMethod', ARGV[6], 'IP', ARGV[7], 'UserAgent', ARGV[8],
|
|
'CreatedAt', ARGV[9], 'LastActiveAt', ARGV[10], 'ExpiresAt', ARGV[11],
|
|
'RevokedAt', ARGV[12], 'RevokedReason', ARGV[13], 'CacheSchema', ARGV[14])
|
|
if ARGV[5] == 'active' then
|
|
redis.call('PEXPIREAT', KEYS[1], ARGV[15])
|
|
else
|
|
redis.call('PEXPIRE', KEYS[1], ARGV[15])
|
|
end
|
|
return 1`
|
|
result, err := common.RDB.Eval(context.Background(), script, []string{userSessionCacheKey(entry.SID)},
|
|
entry.SID, entry.UserID, entry.Version, entry.UserAuthVersion, entry.Status,
|
|
entry.LoginMethod, entry.IP, entry.UserAgent, entry.CreatedAt, entry.LastActiveAt,
|
|
entry.ExpiresAt, entry.RevokedAt, entry.RevokedReason, entry.CacheSchema, redisExpiration,
|
|
).Int()
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if result == 0 {
|
|
return ErrUserSessionInactive
|
|
}
|
|
if entry.Status == UserSessionStatusActive {
|
|
completedAt := time.Now()
|
|
if !completedAt.Before(cacheDeadline) {
|
|
return errUserSessionCacheObservationStale
|
|
}
|
|
if !completedAt.Before(sessionExpiresAt) {
|
|
return ErrUserSessionInactive
|
|
}
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func confirmUserSessionActiveSnapshot(session *UserSession) error {
|
|
if session == nil || session.SID == "" || session.UserID <= 0 || session.Version <= 0 || session.UserAuthVersion <= 0 {
|
|
return ErrUserSessionInvalid
|
|
}
|
|
var count int64
|
|
err := DB.Model(&UserSession{}).
|
|
Where(
|
|
"sid = ? AND user_id = ? AND status = ? AND revoked_at = ? AND expires_at > ? AND version = ? AND user_auth_version = ?",
|
|
session.SID,
|
|
session.UserID,
|
|
UserSessionStatusActive,
|
|
0,
|
|
time.Now().Unix(),
|
|
session.Version,
|
|
session.UserAuthVersion,
|
|
).
|
|
Count(&count).Error
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if count != 1 {
|
|
return ErrUserSessionInactive
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func writeUserSessionDenyFence(session *UserSession, status string, now int64, reason string) error {
|
|
if !common.RedisEnabled {
|
|
return nil
|
|
}
|
|
entry := session.cacheEntry()
|
|
entry.Status = status
|
|
entry.RevokedAt = now
|
|
entry.RevokedReason = reason
|
|
return writeUserSessionCache(entry, time.Time{})
|
|
}
|
|
|
|
func ListActiveUserSessions(userID int, currentSID string, now int64) ([]UserSession, error) {
|
|
if userID <= 0 {
|
|
return nil, ErrUserSessionInvalid
|
|
}
|
|
if now <= 0 {
|
|
now = time.Now().Unix()
|
|
}
|
|
var authVersion int64
|
|
if err := DB.Model(&User{}).Where("id = ?", userID).Select("auth_version").Find(&authVersion).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
if authVersion <= 0 {
|
|
return nil, ErrUserSessionInvalid
|
|
}
|
|
sessions := make([]UserSession, 0, userSessionListLimit)
|
|
if currentSID != "" {
|
|
var current []UserSession
|
|
if err := DB.Where(
|
|
"user_id = ? AND user_auth_version = ? AND status = ? AND expires_at > ? AND sid = ?",
|
|
userID,
|
|
authVersion,
|
|
UserSessionStatusActive,
|
|
now,
|
|
currentSID,
|
|
).Limit(1).Find(¤t).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
if len(current) == 1 {
|
|
sessions = append(sessions, current[0])
|
|
}
|
|
}
|
|
remainingLimit := userSessionListLimit - len(sessions)
|
|
|
|
otherQuery := DB.Where(
|
|
"user_id = ? AND user_auth_version = ? AND status = ? AND expires_at > ?",
|
|
userID,
|
|
authVersion,
|
|
UserSessionStatusActive,
|
|
now,
|
|
)
|
|
if currentSID != "" {
|
|
otherQuery = otherQuery.Where("sid <> ?", currentSID)
|
|
}
|
|
var others []UserSession
|
|
if err := otherQuery.Order("last_active_at DESC").Order("created_at DESC").Limit(remainingLimit).Find(&others).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
sessions = append(sessions, others...)
|
|
return sessions, nil
|
|
}
|
|
|
|
// RotateUserSessionRefresh atomically rotates HMAC digests. The UPDATE itself
|
|
// is a compare-and-swap so SQLite, where lockForUpdate is intentionally a
|
|
// no-op, has the same single-winner behavior as MySQL and PostgreSQL. Only a
|
|
// recognized previous digest outside its grace window is treated as reuse;
|
|
// an unknown secret never revokes the victim session.
|
|
func RotateUserSessionRefresh(userID int, sid, presentedHash, nextHash string, now int64, grace time.Duration) (*UserSession, error) {
|
|
if userID <= 0 || sid == "" || presentedHash == "" || nextHash == "" || hmac.Equal([]byte(presentedHash), []byte(nextHash)) {
|
|
return nil, ErrUserSessionInvalid
|
|
}
|
|
if now <= 0 {
|
|
now = time.Now().Unix()
|
|
}
|
|
graceSeconds := int64(grace / time.Second)
|
|
if graceSeconds < 0 {
|
|
return nil, ErrUserSessionInvalid
|
|
}
|
|
for range 3 {
|
|
cacheDeadline := userSessionCacheDeadline()
|
|
var session UserSession
|
|
if err := DB.Where("sid = ? AND user_id = ?", sid, userID).First(&session).Error; err != nil {
|
|
return nil, err
|
|
}
|
|
if session.Status != UserSessionStatusActive || session.RevokedAt != 0 || session.ExpiresAt <= now {
|
|
return nil, ErrUserSessionInactive
|
|
}
|
|
|
|
if hmac.Equal([]byte(session.RefreshHash), []byte(presentedHash)) {
|
|
result := DB.Model(&UserSession{}).
|
|
Where("sid = ? AND user_id = ? AND status = ? AND revoked_at = ? AND expires_at > ? AND refresh_hash = ?",
|
|
sid, userID, UserSessionStatusActive, 0, now, presentedHash).
|
|
Updates(map[string]interface{}{
|
|
"previous_refresh_hash": session.RefreshHash,
|
|
"previous_valid_until": now + graceSeconds,
|
|
"refresh_hash": nextHash,
|
|
"last_active_at": now,
|
|
})
|
|
if result.Error != nil {
|
|
return nil, result.Error
|
|
}
|
|
if result.RowsAffected == 0 {
|
|
continue
|
|
}
|
|
session.PreviousRefreshHash = session.RefreshHash
|
|
session.PreviousValidUntil = now + graceSeconds
|
|
session.RefreshHash = nextHash
|
|
session.LastActiveAt = now
|
|
if err := writeUserSessionCache(session.cacheEntry(), cacheDeadline); err != nil {
|
|
if errors.Is(err, errUserSessionCacheObservationStale) {
|
|
if confirmErr := confirmUserSessionActiveSnapshot(&session); confirmErr != nil {
|
|
return nil, confirmErr
|
|
}
|
|
} else if errors.Is(err, ErrUserSessionInactive) {
|
|
return nil, err
|
|
} else {
|
|
common.SysLog("failed to update rotated user session cache: " + err.Error())
|
|
}
|
|
}
|
|
return &session, nil
|
|
}
|
|
|
|
if session.PreviousRefreshHash == "" || !hmac.Equal([]byte(session.PreviousRefreshHash), []byte(presentedHash)) {
|
|
return nil, ErrUserSessionRefreshInvalid
|
|
}
|
|
if now <= session.PreviousValidUntil {
|
|
return &session, ErrUserSessionRefreshRace
|
|
}
|
|
|
|
// Once a known previous token is replayed outside the grace window the
|
|
// whole token family is compromised. Publish the deny fence first, then
|
|
// revoke the active row regardless of a concurrent refresh rotation.
|
|
if err := writeUserSessionDenyFence(&session, UserSessionStatusRevoking, now, "refresh_reuse"); err != nil {
|
|
return nil, err
|
|
}
|
|
result := DB.Model(&UserSession{}).
|
|
Where("sid = ? AND user_id = ? AND status = ? AND revoked_at = ? AND expires_at > ?",
|
|
sid, userID, UserSessionStatusActive, 0, now).
|
|
Updates(map[string]interface{}{
|
|
"status": UserSessionStatusRevoked,
|
|
"revoked_at": now,
|
|
"revoked_reason": "refresh_reuse",
|
|
})
|
|
if result.Error != nil {
|
|
return nil, result.Error
|
|
}
|
|
if result.RowsAffected == 0 {
|
|
return nil, ErrUserSessionInactive
|
|
}
|
|
session.Status = UserSessionStatusRevoked
|
|
session.RevokedAt = now
|
|
session.RevokedReason = "refresh_reuse"
|
|
if err := writeUserSessionCache(session.cacheEntry(), time.Time{}); err != nil {
|
|
common.SysLog("failed to cache refresh-reuse session revoke: " + err.Error())
|
|
}
|
|
return nil, ErrUserSessionRefreshReuse
|
|
}
|
|
return nil, ErrUserSessionRefreshInvalid
|
|
}
|
|
|
|
func RevokeUserSession(userID int, sid, reason string) (bool, error) {
|
|
if userID <= 0 || sid == "" {
|
|
return false, ErrUserSessionInvalid
|
|
}
|
|
now := time.Now().Unix()
|
|
var candidate UserSession
|
|
if err := DB.Where("sid = ? AND user_id = ?", sid, userID).First(&candidate).Error; err != nil {
|
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return false, nil
|
|
}
|
|
return false, err
|
|
}
|
|
if candidate.Status != UserSessionStatusActive || candidate.RevokedAt != 0 || candidate.ExpiresAt <= now {
|
|
return false, nil
|
|
}
|
|
if err := writeUserSessionDenyFence(&candidate, UserSessionStatusRevoking, now, reason); err != nil {
|
|
return false, err
|
|
}
|
|
|
|
var revoked bool
|
|
err := DB.Transaction(func(tx *gorm.DB) error {
|
|
var current UserSession
|
|
if err := lockForUpdate(tx).Where("sid = ? AND user_id = ?", sid, userID).First(¤t).Error; err != nil {
|
|
return err
|
|
}
|
|
if current.Status != UserSessionStatusActive || current.RevokedAt != 0 || current.ExpiresAt <= now {
|
|
return nil
|
|
}
|
|
result := tx.Model(&UserSession{}).Where("sid = ? AND status = ?", sid, UserSessionStatusActive).Updates(map[string]interface{}{
|
|
"status": UserSessionStatusRevoked,
|
|
"revoked_at": now,
|
|
"revoked_reason": reason,
|
|
})
|
|
if result.Error != nil {
|
|
return result.Error
|
|
}
|
|
revoked = result.RowsAffected == 1
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
if revoked {
|
|
candidate.Status = UserSessionStatusRevoked
|
|
candidate.RevokedAt = now
|
|
candidate.RevokedReason = reason
|
|
if err := writeUserSessionCache(candidate.cacheEntry(), time.Time{}); err != nil {
|
|
common.SysLog("failed to finalize user session revoke tombstone: " + err.Error())
|
|
}
|
|
}
|
|
return revoked, nil
|
|
}
|
|
|
|
// RevokeUserSessionByRefreshHash is used when logout is authenticated only by
|
|
// the HttpOnly refresh cookie. Possession of a SID alone is insufficient. The
|
|
// immediately previous digest is accepted only inside the refresh race window.
|
|
func RevokeUserSessionByRefreshHash(sid, presentedHash, reason string) (bool, error) {
|
|
if sid == "" || presentedHash == "" {
|
|
return false, ErrUserSessionInvalid
|
|
}
|
|
now := time.Now().Unix()
|
|
var session UserSession
|
|
var revoked bool
|
|
err := DB.Transaction(func(tx *gorm.DB) error {
|
|
if err := lockForUpdate(tx).Where("sid = ?", sid).First(&session).Error; err != nil {
|
|
if errors.Is(err, gorm.ErrRecordNotFound) {
|
|
return nil
|
|
}
|
|
return err
|
|
}
|
|
if session.Status != UserSessionStatusActive || session.RevokedAt != 0 || session.ExpiresAt <= now {
|
|
return nil
|
|
}
|
|
validCurrent := hmac.Equal([]byte(session.RefreshHash), []byte(presentedHash))
|
|
validPrevious := session.PreviousRefreshHash != "" && now <= session.PreviousValidUntil &&
|
|
hmac.Equal([]byte(session.PreviousRefreshHash), []byte(presentedHash))
|
|
if !validCurrent && !validPrevious {
|
|
return nil
|
|
}
|
|
if err := writeUserSessionDenyFence(&session, UserSessionStatusRevoking, now, reason); err != nil {
|
|
return err
|
|
}
|
|
result := tx.Model(&UserSession{}).Where("sid = ? AND status = ?", sid, UserSessionStatusActive).Updates(map[string]interface{}{
|
|
"status": UserSessionStatusRevoked,
|
|
"revoked_at": now,
|
|
"revoked_reason": reason,
|
|
})
|
|
if result.Error != nil {
|
|
return result.Error
|
|
}
|
|
revoked = result.RowsAffected == 1
|
|
if revoked {
|
|
session.Status = UserSessionStatusRevoked
|
|
session.RevokedAt = now
|
|
session.RevokedReason = reason
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
if revoked {
|
|
if err := writeUserSessionCache(session.cacheEntry(), time.Time{}); err != nil {
|
|
common.SysLog("failed to finalize refresh-authenticated session revoke tombstone: " + err.Error())
|
|
}
|
|
}
|
|
return revoked, nil
|
|
}
|
|
|
|
// AdvanceUserSessionAuthVersion preserves one browser session across a
|
|
// user-level security-version change. Both old access JWTs and concurrent
|
|
// updates are invalidated by advancing the per-session version as well.
|
|
func AdvanceUserSessionAuthVersion(userID int, sid string, expectedSessionVersion, expectedUserAuthVersion, nextUserAuthVersion int64) (*UserSession, error) {
|
|
if userID <= 0 || sid == "" || expectedSessionVersion <= 0 || expectedUserAuthVersion <= 0 || nextUserAuthVersion <= expectedUserAuthVersion {
|
|
return nil, ErrUserSessionInvalid
|
|
}
|
|
cacheDeadline := userSessionCacheDeadline()
|
|
now := time.Now().Unix()
|
|
var session UserSession
|
|
err := DB.Transaction(func(tx *gorm.DB) error {
|
|
if err := lockForUpdate(tx).Where("sid = ? AND user_id = ?", sid, userID).First(&session).Error; err != nil {
|
|
return err
|
|
}
|
|
if session.Status != UserSessionStatusActive || session.ExpiresAt <= now ||
|
|
session.Version != expectedSessionVersion || session.UserAuthVersion != expectedUserAuthVersion {
|
|
return ErrUserSessionInactive
|
|
}
|
|
session.Version++
|
|
session.UserAuthVersion = nextUserAuthVersion
|
|
session.LastActiveAt = now
|
|
result := tx.Model(&UserSession{}).
|
|
Where("sid = ? AND status = ? AND version = ? AND user_auth_version = ?", sid, UserSessionStatusActive, expectedSessionVersion, expectedUserAuthVersion).
|
|
Updates(map[string]interface{}{
|
|
"version": session.Version,
|
|
"user_auth_version": session.UserAuthVersion,
|
|
"last_active_at": session.LastActiveAt,
|
|
})
|
|
if result.Error != nil {
|
|
return result.Error
|
|
}
|
|
if result.RowsAffected != 1 {
|
|
return ErrUserSessionInactive
|
|
}
|
|
return nil
|
|
})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if err := writeUserSessionCache(session.cacheEntry(), cacheDeadline); err != nil {
|
|
if errors.Is(err, errUserSessionCacheObservationStale) {
|
|
if confirmErr := confirmUserSessionActiveSnapshot(&session); confirmErr != nil {
|
|
return nil, confirmErr
|
|
}
|
|
} else {
|
|
return nil, err
|
|
}
|
|
}
|
|
return &session, nil
|
|
}
|
|
|
|
func RevokeOtherUserSessions(userID int, currentSID, reason string) (int64, error) {
|
|
return revokeUserSessions(userID, currentSID, reason)
|
|
}
|
|
|
|
func RevokeAllUserSessions(userID int, reason string) (int64, error) {
|
|
return revokeUserSessions(userID, "", reason)
|
|
}
|
|
|
|
func revokeUserSessions(userID int, excludedSID, reason string) (int64, error) {
|
|
if userID <= 0 {
|
|
return 0, ErrUserSessionInvalid
|
|
}
|
|
now := time.Now().Unix()
|
|
var totalAffected int64
|
|
for {
|
|
query := DB.Where("user_id = ? AND status = ? AND expires_at > ?", userID, UserSessionStatusActive, now)
|
|
if excludedSID != "" {
|
|
query = query.Where("sid <> ?", excludedSID)
|
|
}
|
|
var candidates []UserSession
|
|
if err := query.Order("sid").Limit(userSessionRevokeBatchSize).Find(&candidates).Error; err != nil {
|
|
return totalAffected, err
|
|
}
|
|
if len(candidates) == 0 {
|
|
return totalAffected, nil
|
|
}
|
|
for i := range candidates {
|
|
if err := writeUserSessionDenyFence(&candidates[i], UserSessionStatusRevoking, now, reason); err != nil {
|
|
return totalAffected, err
|
|
}
|
|
}
|
|
|
|
sids := make([]string, 0, len(candidates))
|
|
for i := range candidates {
|
|
sids = append(sids, candidates[i].SID)
|
|
}
|
|
var affected int64
|
|
var revoked []UserSession
|
|
err := DB.Transaction(func(tx *gorm.DB) error {
|
|
if err := lockForUpdate(tx).Where("sid IN ? AND status = ?", sids, UserSessionStatusActive).Find(&revoked).Error; err != nil {
|
|
return err
|
|
}
|
|
if len(revoked) == 0 {
|
|
return nil
|
|
}
|
|
lockedSIDs := make([]string, 0, len(revoked))
|
|
for i := range revoked {
|
|
lockedSIDs = append(lockedSIDs, revoked[i].SID)
|
|
}
|
|
result := tx.Model(&UserSession{}).Where("sid IN ? AND status = ?", lockedSIDs, UserSessionStatusActive).Updates(map[string]interface{}{
|
|
"status": UserSessionStatusRevoked,
|
|
"revoked_at": now,
|
|
"revoked_reason": reason,
|
|
})
|
|
affected = result.RowsAffected
|
|
return result.Error
|
|
})
|
|
if err != nil {
|
|
return totalAffected, err
|
|
}
|
|
totalAffected += affected
|
|
for i := range revoked {
|
|
revoked[i].Status = UserSessionStatusRevoked
|
|
revoked[i].RevokedAt = now
|
|
revoked[i].RevokedReason = reason
|
|
if err := writeUserSessionCache(revoked[i].cacheEntry(), time.Time{}); err != nil {
|
|
common.SysLog("failed to finalize bulk user session revoke tombstone: " + err.Error())
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
func DeleteExpiredUserSessions(now int64) error {
|
|
if now <= 0 {
|
|
now = time.Now().Unix()
|
|
}
|
|
if common.UserSessionRevokedRetentionDays <= 0 || common.UserSessionIssuanceWindowSeconds <= 0 {
|
|
return ErrUserSessionInvalid
|
|
}
|
|
issuanceCutoff := now - common.UserSessionIssuanceWindowSeconds
|
|
revokedBefore := now - int64(common.UserSessionRevokedRetentionDays)*24*60*60
|
|
return deleteExpiredUserSessionsBefore(now, issuanceCutoff, revokedBefore)
|
|
}
|
|
|
|
func DeleteOldRevokedUserSessions(now int64) error {
|
|
if now <= 0 {
|
|
now = time.Now().Unix()
|
|
}
|
|
if common.UserSessionRevokedRetentionDays <= 0 || common.UserSessionIssuanceWindowSeconds <= 0 {
|
|
return ErrUserSessionInvalid
|
|
}
|
|
issuanceCutoff := now - common.UserSessionIssuanceWindowSeconds
|
|
revokedBefore := now - int64(common.UserSessionRevokedRetentionDays)*24*60*60
|
|
return deleteRevokedUserSessionsBefore(revokedBefore, issuanceCutoff)
|
|
}
|
|
|
|
func deleteExpiredUserSessionsBefore(expiredBefore, issuanceCutoff, revokedBefore int64) error {
|
|
for {
|
|
var sids []string
|
|
if err := DB.Model(&UserSession{}).
|
|
Where(
|
|
"expires_at < ? AND created_at <= ? AND (status <> ? OR revoked_at <= 0 OR revoked_at < ?)",
|
|
expiredBefore,
|
|
issuanceCutoff,
|
|
UserSessionStatusRevoked,
|
|
revokedBefore,
|
|
).
|
|
Order("expires_at").Limit(userSessionCleanupScanLimit).Pluck("sid", &sids).Error; err != nil {
|
|
return err
|
|
}
|
|
if len(sids) == 0 {
|
|
return nil
|
|
}
|
|
for start := 0; start < len(sids); start += userSessionCleanupBatchSize {
|
|
end := start + userSessionCleanupBatchSize
|
|
if end > len(sids) {
|
|
end = len(sids)
|
|
}
|
|
if err := DB.Where("sid IN ?", sids[start:end]).
|
|
Where(
|
|
"expires_at < ? AND created_at <= ? AND (status <> ? OR revoked_at <= 0 OR revoked_at < ?)",
|
|
expiredBefore,
|
|
issuanceCutoff,
|
|
UserSessionStatusRevoked,
|
|
revokedBefore,
|
|
).
|
|
Delete(&UserSession{}).Error; err != nil {
|
|
return err
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
func deleteRevokedUserSessionsBefore(revokedBefore, issuanceCutoff int64) error {
|
|
for {
|
|
var sids []string
|
|
if err := DB.Model(&UserSession{}).
|
|
Where(
|
|
"status = ? AND revoked_at > 0 AND revoked_at < ? AND created_at <= ?",
|
|
UserSessionStatusRevoked,
|
|
revokedBefore,
|
|
issuanceCutoff,
|
|
).
|
|
Order("revoked_at").Limit(userSessionCleanupScanLimit).Pluck("sid", &sids).Error; err != nil {
|
|
return err
|
|
}
|
|
if len(sids) == 0 {
|
|
return nil
|
|
}
|
|
for start := 0; start < len(sids); start += userSessionCleanupBatchSize {
|
|
end := start + userSessionCleanupBatchSize
|
|
if end > len(sids) {
|
|
end = len(sids)
|
|
}
|
|
if err := DB.Where("sid IN ?", sids[start:end]).
|
|
Where(
|
|
"status = ? AND revoked_at > 0 AND revoked_at < ? AND created_at <= ?",
|
|
UserSessionStatusRevoked,
|
|
revokedBefore,
|
|
issuanceCutoff,
|
|
).
|
|
Delete(&UserSession{}).Error; err != nil {
|
|
return err
|
|
}
|
|
}
|
|
}
|
|
}
|