* 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
432 lines
17 KiB
Go
432 lines
17 KiB
Go
package service
|
|
|
|
import (
|
|
"bytes"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/QuantumNous/new-api/common"
|
|
"github.com/QuantumNous/new-api/model"
|
|
"github.com/alicebob/miniredis/v2"
|
|
"github.com/gin-gonic/gin"
|
|
"github.com/glebarez/sqlite"
|
|
"github.com/go-redis/redis/v8"
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
"gorm.io/gorm"
|
|
)
|
|
|
|
func setupAuthSessionTestDB(t *testing.T) *model.User {
|
|
t.Helper()
|
|
previousDB, previousRedis := model.DB, common.RedisEnabled
|
|
previousActiveLimit := common.UserSessionActiveLimit
|
|
previousIssuanceLimit := common.UserSessionIssuanceLimit
|
|
previousIssuanceWindow := common.UserSessionIssuanceWindowSeconds
|
|
previousRevokedRetention := common.UserSessionRevokedRetentionDays
|
|
previousAlertThreshold := common.UserSessionHourlyAlertThreshold
|
|
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
|
require.NoError(t, err)
|
|
sqlDB, err := db.DB()
|
|
require.NoError(t, err)
|
|
sqlDB.SetMaxOpenConns(1)
|
|
require.NoError(t, db.AutoMigrate(&model.User{}, &model.UserSession{}, &model.AuthFlow{}))
|
|
model.DB = db
|
|
common.RedisEnabled = false
|
|
common.UserSessionActiveLimit = common.DefaultUserSessionActiveLimit
|
|
common.UserSessionIssuanceLimit = common.DefaultUserSessionIssuanceLimit
|
|
common.UserSessionIssuanceWindowSeconds = int64(common.DefaultUserSessionIssuanceWindowSeconds)
|
|
common.UserSessionRevokedRetentionDays = common.DefaultUserSessionRevokedRetentionDays
|
|
common.UserSessionHourlyAlertThreshold = common.DefaultUserSessionHourlyAlertThreshold
|
|
t.Cleanup(func() {
|
|
model.DB = previousDB
|
|
common.RedisEnabled = previousRedis
|
|
common.UserSessionActiveLimit = previousActiveLimit
|
|
common.UserSessionIssuanceLimit = previousIssuanceLimit
|
|
common.UserSessionIssuanceWindowSeconds = previousIssuanceWindow
|
|
common.UserSessionRevokedRetentionDays = previousRevokedRetention
|
|
common.UserSessionHourlyAlertThreshold = previousAlertThreshold
|
|
_ = sqlDB.Close()
|
|
})
|
|
user := &model.User{
|
|
Username: "session-user",
|
|
Password: "unused-password-hash",
|
|
Role: common.RoleCommonUser,
|
|
Status: common.UserStatusEnabled,
|
|
Group: "default",
|
|
AuthVersion: 1,
|
|
}
|
|
require.NoError(t, db.Create(user).Error)
|
|
return user
|
|
}
|
|
|
|
func useIndependentAuthSessionRedis(t *testing.T) (*miniredis.Miniredis, *redis.Client, *miniredis.Miniredis, *redis.Client) {
|
|
t.Helper()
|
|
previousRedisEnabled := common.RedisEnabled
|
|
previousRDB := common.RDB
|
|
previousSyncFrequency := common.SyncFrequency
|
|
serverA := miniredis.RunT(t)
|
|
serverB := miniredis.RunT(t)
|
|
clientA := redis.NewClient(&redis.Options{Addr: serverA.Addr()})
|
|
clientB := redis.NewClient(&redis.Options{Addr: serverB.Addr()})
|
|
common.RedisEnabled = true
|
|
common.SyncFrequency = 2
|
|
common.RDB = clientA
|
|
t.Cleanup(func() {
|
|
_ = clientA.Close()
|
|
_ = clientB.Close()
|
|
common.RedisEnabled = previousRedisEnabled
|
|
common.RDB = previousRDB
|
|
common.SyncFrequency = previousSyncFrequency
|
|
})
|
|
return serverA, clientA, serverB, clientB
|
|
}
|
|
|
|
func cachedLoginSessionKey(t *testing.T, server *miniredis.Miniredis) string {
|
|
t.Helper()
|
|
for _, key := range server.Keys() {
|
|
if strings.HasPrefix(key, "auth:session:") {
|
|
return key
|
|
}
|
|
}
|
|
require.FailNow(t, "login session was not cached")
|
|
return ""
|
|
}
|
|
|
|
func TestCreateLoginSessionEnforcesActiveLimitAcrossAuthVersions(t *testing.T) {
|
|
useTestSessionSecret(t)
|
|
user := setupAuthSessionTestDB(t)
|
|
common.UserSessionActiveLimit = 50
|
|
common.UserSessionIssuanceLimit = 100
|
|
now := time.Now().Unix()
|
|
rows := make([]model.UserSession, 0, 49)
|
|
for i := 0; i < 49; i++ {
|
|
authVersion := user.AuthVersion
|
|
if i == 0 {
|
|
authVersion++
|
|
}
|
|
rows = append(rows, model.UserSession{
|
|
SID: fmt.Sprintf("active-limit-%02d", i),
|
|
UserID: user.Id,
|
|
Version: 1,
|
|
UserAuthVersion: authVersion,
|
|
Status: model.UserSessionStatusActive,
|
|
RefreshHash: fmt.Sprintf("hash-%02d", i),
|
|
LoginMethod: "password",
|
|
CreatedAt: now - int64(i),
|
|
LastActiveAt: now - int64(i),
|
|
ExpiresAt: now + 3600,
|
|
})
|
|
}
|
|
require.NoError(t, model.DB.Create(&rows).Error)
|
|
|
|
_, err := CreateLoginSession(user.Id, "password", "127.0.0.1", "test-agent")
|
|
require.NoError(t, err, "49 active sessions must allow creation of the 50th")
|
|
|
|
_, err = CreateLoginSession(user.Id, "password", "127.0.0.1", "test-agent")
|
|
assert.ErrorIs(t, err, model.ErrUserSessionLimit)
|
|
var count int64
|
|
require.NoError(t, model.DB.Model(&model.UserSession{}).Count(&count).Error)
|
|
assert.Equal(t, int64(50), count)
|
|
}
|
|
|
|
func TestCreateLoginSessionEnforcesIssuanceLimitAcrossAllStatuses(t *testing.T) {
|
|
useTestSessionSecret(t)
|
|
user := setupAuthSessionTestDB(t)
|
|
common.UserSessionActiveLimit = 10
|
|
common.UserSessionIssuanceLimit = 3
|
|
common.UserSessionIssuanceWindowSeconds = 60
|
|
now := time.Now().Unix()
|
|
rows := []model.UserSession{
|
|
{
|
|
SID: "issuance-limit-revoked", UserID: user.Id, Version: 1, UserAuthVersion: user.AuthVersion + 1,
|
|
Status: model.UserSessionStatusRevoked, RefreshHash: "hash-revoked", LoginMethod: "password",
|
|
CreatedAt: now - 2, LastActiveAt: now - 2, ExpiresAt: now + 3600, RevokedAt: now - 1,
|
|
},
|
|
{
|
|
SID: "issuance-limit-expired", UserID: user.Id, Version: 1, UserAuthVersion: user.AuthVersion,
|
|
Status: model.UserSessionStatusActive, RefreshHash: "hash-expired", LoginMethod: "password",
|
|
CreatedAt: now - 1, LastActiveAt: now - 1, ExpiresAt: now - 1,
|
|
},
|
|
{
|
|
SID: "issuance-outside-effective-window", UserID: user.Id, Version: 1, UserAuthVersion: user.AuthVersion,
|
|
Status: model.UserSessionStatusRevoked, RefreshHash: "hash-outside", LoginMethod: "password",
|
|
CreatedAt: now - 61, LastActiveAt: now - 61, ExpiresAt: now + 3600, RevokedAt: now - 60,
|
|
},
|
|
}
|
|
require.NoError(t, model.DB.Create(&rows).Error)
|
|
|
|
_, err := CreateLoginSession(user.Id, "password", "127.0.0.1", "test-agent")
|
|
require.NoError(t, err, "rows outside the effective issuance window must not consume the limit")
|
|
|
|
_, err = CreateLoginSession(user.Id, "password", "127.0.0.1", "test-agent")
|
|
assert.ErrorIs(t, err, model.ErrUserSessionIssuanceLimit)
|
|
var count int64
|
|
require.NoError(t, model.DB.Model(&model.UserSession{}).Count(&count).Error)
|
|
assert.Equal(t, int64(4), count)
|
|
}
|
|
|
|
func TestPasswordResetDoesNotClearSessionIssuanceHistory(t *testing.T) {
|
|
useTestSessionSecret(t)
|
|
user := setupAuthSessionTestDB(t)
|
|
common.UserSessionActiveLimit = 50
|
|
common.UserSessionIssuanceLimit = 1
|
|
email := "session-reset@example.com"
|
|
require.NoError(t, model.DB.Model(user).Update("email", email).Error)
|
|
|
|
_, err := CreateLoginSession(user.Id, "password", "127.0.0.1", "test-agent")
|
|
require.NoError(t, err)
|
|
require.NoError(t, model.ResetUserPasswordByEmail(email, "new-password"))
|
|
|
|
_, err = CreateLoginSession(user.Id, "password", "127.0.0.1", "test-agent")
|
|
assert.ErrorIs(t, err, model.ErrUserSessionIssuanceLimit)
|
|
}
|
|
|
|
func TestCreateLoginSessionFailsClosedWhenLimitCountFails(t *testing.T) {
|
|
useTestSessionSecret(t)
|
|
user := setupAuthSessionTestDB(t)
|
|
forcedErr := errors.New("forced session count failure")
|
|
callbackName := "test:fail_user_session_limit_count"
|
|
callbackRegistered := true
|
|
require.NoError(t, model.DB.Callback().Query().Before("gorm:query").Register(callbackName, func(tx *gorm.DB) {
|
|
if tx.Statement != nil && tx.Statement.Table == "user_sessions" {
|
|
tx.AddError(forcedErr)
|
|
}
|
|
}))
|
|
t.Cleanup(func() {
|
|
if callbackRegistered {
|
|
_ = model.DB.Callback().Query().Remove(callbackName)
|
|
}
|
|
})
|
|
|
|
_, err := CreateLoginSession(user.Id, "password", "127.0.0.1", "test-agent")
|
|
assert.ErrorIs(t, err, forcedErr)
|
|
require.NoError(t, model.DB.Callback().Query().Remove(callbackName))
|
|
callbackRegistered = false
|
|
var count int64
|
|
require.NoError(t, model.DB.Model(&model.UserSession{}).Count(&count).Error)
|
|
assert.Zero(t, count)
|
|
}
|
|
|
|
func TestCleanupAuthArtifactsAlertsBeforeDeletingHourlyIssuance(t *testing.T) {
|
|
setupAuthSessionTestDB(t)
|
|
common.UserSessionHourlyAlertThreshold = 2
|
|
common.UserSessionIssuanceWindowSeconds = 1
|
|
now := time.Now()
|
|
boundaryRows := make([]model.UserSession, 0, 2)
|
|
for i := 0; i < 2; i++ {
|
|
boundaryRows = append(boundaryRows, model.UserSession{
|
|
SID: "hourly-boundary-" + string(rune('a'+i)), UserID: 1, Version: 1, UserAuthVersion: 1,
|
|
Status: model.UserSessionStatusActive, RefreshHash: "hash", LoginMethod: "password",
|
|
CreatedAt: now.Add(-2 * time.Second).Unix(), LastActiveAt: now.Add(-time.Hour).Unix(), ExpiresAt: now.Add(-time.Minute).Unix(),
|
|
})
|
|
}
|
|
require.NoError(t, model.DB.Create(&boundaryRows).Error)
|
|
|
|
var logBuffer bytes.Buffer
|
|
common.LogWriterMu.Lock()
|
|
previousErrorWriter := gin.DefaultErrorWriter
|
|
gin.DefaultErrorWriter = &logBuffer
|
|
common.LogWriterMu.Unlock()
|
|
t.Cleanup(func() {
|
|
common.LogWriterMu.Lock()
|
|
gin.DefaultErrorWriter = previousErrorWriter
|
|
common.LogWriterMu.Unlock()
|
|
})
|
|
|
|
cleanupAuthArtifacts()
|
|
assert.Empty(t, logBuffer.String(), "the hourly alert uses a strict greater-than threshold")
|
|
var count int64
|
|
require.NoError(t, model.DB.Model(&model.UserSession{}).Count(&count).Error)
|
|
assert.Zero(t, count)
|
|
|
|
exceededRows := make([]model.UserSession, 0, 3)
|
|
for i := 0; i < 3; i++ {
|
|
exceededRows = append(exceededRows, model.UserSession{
|
|
SID: "hourly-exceeded-" + string(rune('a'+i)), UserID: 1, Version: 1, UserAuthVersion: 1,
|
|
Status: model.UserSessionStatusActive, RefreshHash: "hash", LoginMethod: "password",
|
|
CreatedAt: now.Add(-2 * time.Second).Unix(), LastActiveAt: now.Add(-time.Hour).Unix(), ExpiresAt: now.Add(-time.Minute).Unix(),
|
|
})
|
|
}
|
|
require.NoError(t, model.DB.Create(&exceededRows).Error)
|
|
logBuffer.Reset()
|
|
cleanupAuthArtifacts()
|
|
assert.Contains(t, logBuffer.String(), "hourly user session issuance exceeded alert threshold")
|
|
require.NoError(t, model.DB.Model(&model.UserSession{}).Count(&count).Error)
|
|
assert.Zero(t, count, "alerting must happen before expired rows are deleted")
|
|
}
|
|
|
|
func TestCleanupAuthArtifactsRemovesOnlyExpiredRecords(t *testing.T) {
|
|
setupAuthSessionTestDB(t)
|
|
now := time.Now()
|
|
oldExpiry := now.Add(-25 * time.Hour)
|
|
require.NoError(t, model.DB.Create(&model.UserSession{
|
|
SID: "expired-session", UserID: 1, Version: 1, UserAuthVersion: 1,
|
|
Status: model.UserSessionStatusActive, RefreshHash: "hash", LoginMethod: "password",
|
|
CreatedAt: oldExpiry.Unix(), LastActiveAt: oldExpiry.Unix(), ExpiresAt: oldExpiry.Unix(),
|
|
}).Error)
|
|
require.NoError(t, model.DB.Create(&model.AuthFlow{
|
|
TokenHash: "expired-flow", Purpose: model.AuthFlowPurposeTwoFALogin,
|
|
ExpiresAt: oldExpiry,
|
|
}).Error)
|
|
require.NoError(t, model.DB.Create(&model.AuthFlow{
|
|
TokenHash: "recent-flow", Purpose: model.AuthFlowPurposeTwoFALogin,
|
|
ExpiresAt: now.Add(time.Minute),
|
|
}).Error)
|
|
|
|
cleanupAuthArtifacts()
|
|
|
|
var sessionCount int64
|
|
require.NoError(t, model.DB.Model(&model.UserSession{}).Count(&sessionCount).Error)
|
|
assert.Zero(t, sessionCount)
|
|
var flows []model.AuthFlow
|
|
require.NoError(t, model.DB.Find(&flows).Error)
|
|
require.Len(t, flows, 1)
|
|
assert.Equal(t, "recent-flow", flows[0].TokenHash)
|
|
}
|
|
|
|
func TestCleanupAuthArtifactsContinuesWithRevokedCleanupAfterExpiredBatchFailure(t *testing.T) {
|
|
setupAuthSessionTestDB(t)
|
|
now := time.Now()
|
|
oldCreatedAt := now.Add(-8 * 24 * time.Hour).Unix()
|
|
require.NoError(t, model.DB.Create(&[]model.UserSession{
|
|
{
|
|
SID: "failed-expired-cleanup", UserID: 1, Version: 1, UserAuthVersion: 1,
|
|
Status: model.UserSessionStatusActive, RefreshHash: "hash-expired", LoginMethod: "password",
|
|
CreatedAt: oldCreatedAt, LastActiveAt: oldCreatedAt, ExpiresAt: now.Add(-time.Minute).Unix(),
|
|
},
|
|
{
|
|
SID: "independent-revoked-cleanup", UserID: 1, Version: 1, UserAuthVersion: 1,
|
|
Status: model.UserSessionStatusRevoked, RefreshHash: "hash-revoked", LoginMethod: "password",
|
|
CreatedAt: oldCreatedAt, LastActiveAt: oldCreatedAt, ExpiresAt: now.Add(time.Hour).Unix(),
|
|
RevokedAt: now.Add(-8 * 24 * time.Hour).Unix(),
|
|
},
|
|
}).Error)
|
|
|
|
forcedErr := errors.New("forced expired cleanup failure")
|
|
callbackName := "test:fail_first_user_session_cleanup_batch"
|
|
failedFirstDelete := false
|
|
require.NoError(t, model.DB.Callback().Delete().Before("gorm:delete").Register(callbackName, func(tx *gorm.DB) {
|
|
if tx.Statement != nil && tx.Statement.Table == "user_sessions" && !failedFirstDelete {
|
|
failedFirstDelete = true
|
|
tx.AddError(forcedErr)
|
|
}
|
|
}))
|
|
t.Cleanup(func() { _ = model.DB.Callback().Delete().Remove(callbackName) })
|
|
|
|
cleanupAuthArtifacts()
|
|
|
|
var expired model.UserSession
|
|
require.NoError(t, model.DB.First(&expired, "sid = ?", "failed-expired-cleanup").Error)
|
|
var revoked model.UserSession
|
|
assert.ErrorIs(t, model.DB.First(&revoked, "sid = ?", "independent-revoked-cleanup").Error, gorm.ErrRecordNotFound)
|
|
}
|
|
|
|
func TestLoginSessionCreateRefreshAndRevoke(t *testing.T) {
|
|
useTestSessionSecret(t)
|
|
user := setupAuthSessionTestDB(t)
|
|
|
|
bundle, err := CreateLoginSession(user.Id, "password", "127.0.0.1", "test-agent")
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, bundle.RefreshToken)
|
|
identity, err := ParseAccessToken(bundle.AccessToken)
|
|
require.NoError(t, err)
|
|
_, cachedUser, err := ValidateLoginSession(identity)
|
|
require.NoError(t, err)
|
|
assert.Equal(t, user.Id, cachedUser.Id)
|
|
require.NoError(t, RevokeByRefreshToken(bundle.Session.SID+".wrong-refresh-secret", "", "logout"))
|
|
_, _, err = ValidateLoginSession(identity)
|
|
require.NoError(t, err, "a caller that only knows sid must not be able to revoke the session")
|
|
|
|
refreshed, _, err := RefreshLoginSession(bundle.RefreshToken, bundle.Session.SID, "127.0.0.2", "test-agent-2")
|
|
require.NoError(t, err)
|
|
assert.NotEqual(t, bundle.RefreshToken, refreshed.RefreshToken)
|
|
recovered, _, err := RefreshLoginSession(bundle.RefreshToken, bundle.Session.SID, "127.0.0.2", "test-agent-2")
|
|
require.NoError(t, err)
|
|
assert.Equal(t, refreshed.RefreshToken, recovered.RefreshToken, "a concurrent refresh must recover the winner's rotated token")
|
|
|
|
_, _, err = RefreshLoginSession(refreshed.RefreshToken, "different-session", "127.0.0.2", "test-agent-2")
|
|
assert.ErrorIs(t, err, ErrLoginSessionMismatch)
|
|
|
|
require.NoError(t, RevokeByRefreshToken(refreshed.RefreshToken, refreshed.Session.SID, "logout"))
|
|
_, _, err = ValidateLoginSession(identity)
|
|
assert.True(t, errors.Is(err, ErrLoginSessionRevoked))
|
|
}
|
|
|
|
func TestIndependentRedisSessionRevokeConvergesAfterCacheTTL(t *testing.T) {
|
|
useTestSessionSecret(t)
|
|
user := setupAuthSessionTestDB(t)
|
|
_, clientA, serverB, clientB := useIndependentAuthSessionRedis(t)
|
|
|
|
common.RDB = clientA
|
|
bundle, err := CreateLoginSession(user.Id, "password", "127.0.0.1", "node-a")
|
|
require.NoError(t, err)
|
|
identity, err := ParseAccessToken(bundle.AccessToken)
|
|
require.NoError(t, err)
|
|
|
|
common.RDB = clientB
|
|
_, _, err = ValidateLoginSession(identity)
|
|
require.NoError(t, err)
|
|
assert.NotEmpty(t, cachedLoginSessionKey(t, serverB), "node B must hold its own session cache entry")
|
|
|
|
common.RDB = clientA
|
|
require.NoError(t, RevokeByRefreshToken(bundle.RefreshToken, bundle.Session.SID, "logout"))
|
|
|
|
serverB.FastForward(3 * time.Second)
|
|
common.RDB = clientB
|
|
_, _, err = ValidateLoginSession(identity)
|
|
assert.ErrorIs(t, err, ErrLoginSessionRevoked)
|
|
}
|
|
|
|
func TestIndependentRedisAuthVersionAdvanceConvergesAfterCacheTTL(t *testing.T) {
|
|
useTestSessionSecret(t)
|
|
user := setupAuthSessionTestDB(t)
|
|
_, clientA, serverB, clientB := useIndependentAuthSessionRedis(t)
|
|
|
|
common.RDB = clientA
|
|
bundle, err := CreateLoginSession(user.Id, "password", "127.0.0.1", "node-a")
|
|
require.NoError(t, err)
|
|
oldIdentity, err := ParseAccessToken(bundle.AccessToken)
|
|
require.NoError(t, err)
|
|
|
|
common.RDB = clientB
|
|
_, _, err = ValidateLoginSession(oldIdentity)
|
|
require.NoError(t, err)
|
|
cacheKey := cachedLoginSessionKey(t, serverB)
|
|
version := serverB.HGet(cacheKey, "Version")
|
|
assert.Equal(t, "1", version, "node B must hold the pre-rotation session version")
|
|
|
|
common.RDB = clientA
|
|
rotated, err := AdvanceCurrentSessionSecurity(oldIdentity, "security_update")
|
|
require.NoError(t, err)
|
|
newIdentity, err := ParseAccessToken(rotated.AccessToken)
|
|
require.NoError(t, err)
|
|
assert.Greater(t, newIdentity.SessionVersion, oldIdentity.SessionVersion)
|
|
assert.Greater(t, newIdentity.UserAuthVersion, oldIdentity.UserAuthVersion)
|
|
|
|
serverB.FastForward(3 * time.Second)
|
|
common.RDB = clientB
|
|
_, _, err = ValidateLoginSession(newIdentity)
|
|
require.NoError(t, err)
|
|
_, _, err = ValidateLoginSession(oldIdentity)
|
|
assert.ErrorIs(t, err, ErrLoginSessionRevoked)
|
|
}
|
|
|
|
func TestUserAuthVersionInvalidatesExistingSession(t *testing.T) {
|
|
useTestSessionSecret(t)
|
|
user := setupAuthSessionTestDB(t)
|
|
bundle, err := CreateLoginSession(user.Id, "password", "127.0.0.1", "test-agent")
|
|
require.NoError(t, err)
|
|
identity, err := ParseAccessToken(bundle.AccessToken)
|
|
require.NoError(t, err)
|
|
|
|
_, err = model.BumpUserAuthVersion(user.Id)
|
|
require.NoError(t, err)
|
|
_, _, err = ValidateLoginSession(identity)
|
|
assert.ErrorIs(t, err, ErrLoginSessionRevoked)
|
|
_, err = CreateLoginSessionAtAuthVersion(user.Id, identity.UserAuthVersion, "2fa", "127.0.0.1", "test-agent")
|
|
assert.ErrorIs(t, err, ErrLoginSessionRevoked, "a pending 2FA flow must not survive an auth-version change")
|
|
}
|