refactor(auth): replace dashboard sessions with stateless tokens and session control (#6329)

* 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
This commit is contained in:
Calcium-Ion
2026-07-20 16:48:43 +08:00
committed by GitHub
parent 5a6c53d496
commit 31d70fca39
1605 changed files with 17511 additions and 147913 deletions
+51
View File
@@ -0,0 +1,51 @@
package service
import (
"fmt"
"time"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
)
const authArtifactCleanupInterval = time.Hour
// StartAuthArtifactCleanup removes expired dashboard Sessions and old
// one-time authentication flows. Only the master instance performs cleanup.
func StartAuthArtifactCleanup() {
if !common.IsMasterNode {
return
}
go func() {
cleanupAuthArtifacts()
ticker := time.NewTicker(authArtifactCleanupInterval)
defer ticker.Stop()
for range ticker.C {
cleanupAuthArtifacts()
}
}()
}
func cleanupAuthArtifacts() {
now := time.Now()
count, err := model.CountUserSessionsCreatedSince(0, now.Add(-time.Hour).Unix())
if err != nil {
common.SysError("failed to count hourly user session issuance: " + err.Error())
} else if count > int64(common.UserSessionHourlyAlertThreshold) {
common.SysError(fmt.Sprintf(
"hourly user session issuance exceeded alert threshold: count=%d threshold=%d window_seconds=%d",
count,
common.UserSessionHourlyAlertThreshold,
int64(time.Hour/time.Second),
))
}
if err := model.DeleteExpiredUserSessions(now.Unix()); err != nil {
common.SysError("failed to delete expired user sessions: " + err.Error())
}
if err := model.DeleteOldRevokedUserSessions(now.Unix()); err != nil {
common.SysError("failed to delete old revoked user sessions: " + err.Error())
}
if err := model.DeleteExpiredAuthFlows(now); err != nil {
common.SysError("failed to delete expired authentication flows: " + err.Error())
}
}
+418
View File
@@ -0,0 +1,418 @@
package service
import (
"errors"
"fmt"
"net/http"
"strings"
"time"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
)
const RefreshCookieName = "new_api_refresh"
var (
ErrLoginSessionInvalid = errors.New("login session is invalid")
ErrLoginSessionRevoked = errors.New("login session is revoked")
ErrLoginSessionMismatch = errors.New("login session does not match the expected session")
ErrRefreshTokenInvalid = errors.New("refresh token is invalid")
ErrRefreshRace = errors.New("refresh token was already rotated")
)
type LoginSessionView struct {
SID string `json:"sid"`
Current bool `json:"current"`
LoginMethod string `json:"login_method"`
IP string `json:"ip"`
UserAgent string `json:"user_agent"`
CreatedAt int64 `json:"created_at"`
LastActiveAt int64 `json:"last_active_at"`
ExpiresAt int64 `json:"expires_at"`
}
type AuthBundle struct {
AccessToken string `json:"access_token"`
TokenType string `json:"token_type"`
AccessExpiresAt int64 `json:"access_expires_at"`
Session LoginSessionView `json:"session"`
RefreshToken string `json:"-"`
}
func CreateLoginSession(userID int, loginMethod, ip, userAgent string) (*AuthBundle, error) {
return createLoginSession(userID, 0, loginMethod, ip, userAgent)
}
func CreateLoginSessionAtAuthVersion(userID int, expectedAuthVersion int64, loginMethod, ip, userAgent string) (*AuthBundle, error) {
if expectedAuthVersion <= 0 {
return nil, ErrLoginSessionInvalid
}
return createLoginSession(userID, expectedAuthVersion, loginMethod, ip, userAgent)
}
func createLoginSession(userID int, expectedAuthVersion int64, loginMethod, ip, userAgent string) (*AuthBundle, error) {
user, err := model.GetUserCache(userID)
if err != nil {
return nil, err
}
if user.Status != common.UserStatusEnabled || user.AuthVersion <= 0 {
return nil, ErrLoginSessionInvalid
}
if expectedAuthVersion > 0 && user.AuthVersion != expectedAuthVersion {
return nil, ErrLoginSessionRevoked
}
now := time.Now().Unix()
activeCount, err := model.CountActiveUserSessions(userID, now)
if err != nil {
return nil, err
}
if activeCount >= int64(common.UserSessionActiveLimit) {
return nil, model.ErrUserSessionLimit
}
issuanceCount, err := model.CountUserSessionsCreatedSince(userID, now-common.UserSessionIssuanceWindowSeconds)
if err != nil {
return nil, err
}
if issuanceCount >= int64(common.UserSessionIssuanceLimit) {
return nil, model.ErrUserSessionIssuanceLimit
}
refreshSecret, err := common.GenerateRandomCharsKey(64)
if err != nil {
return nil, err
}
session := &model.UserSession{
SID: uuid.NewString(),
UserID: userID,
Version: 1,
UserAuthVersion: user.AuthVersion,
Status: model.UserSessionStatusActive,
RefreshHash: hashRefreshSecret(refreshSecret),
LoginMethod: strings.TrimSpace(loginMethod),
IP: truncateAuthMetadata(ip, 64),
UserAgent: truncateAuthMetadata(userAgent, 512),
CreatedAt: now,
LastActiveAt: now,
ExpiresAt: time.Unix(now, 0).Add(LoginSessionTTL).Unix(),
}
if session.LoginMethod == "" {
session.LoginMethod = "unknown"
}
if err := model.CreateUserSession(session); err != nil {
return nil, err
}
bundle, err := issueAuthBundle(session, session.SID+"."+refreshSecret, true)
if err != nil {
_, _ = model.RevokeUserSession(userID, session.SID, "token_issue_failed")
return nil, err
}
return bundle, nil
}
func ValidateLoginSession(identity AuthIdentity) (*model.UserSession, *model.UserBase, error) {
session, err := model.GetUserSessionCached(identity.SessionID)
if err != nil {
if errors.Is(err, model.ErrUserSessionInactive) {
return nil, nil, ErrLoginSessionRevoked
}
return nil, nil, err
}
now := time.Now().Unix()
if session.UserID != identity.UserID || session.Status != model.UserSessionStatusActive || session.RevokedAt != 0 || session.ExpiresAt <= now || session.Version != identity.SessionVersion || session.UserAuthVersion != identity.UserAuthVersion {
return nil, nil, ErrLoginSessionRevoked
}
user, err := model.GetUserCache(identity.UserID)
if err != nil {
return nil, nil, err
}
if user.Status != common.UserStatusEnabled || user.AuthVersion != identity.UserAuthVersion {
return nil, nil, ErrLoginSessionRevoked
}
return session, user, nil
}
// ValidateSessionReference validates a server-side flow bound to an existing
// dashboard session without requiring an access token on the callback request.
func ValidateSessionReference(userID int, sid string) (AuthIdentity, error) {
if userID <= 0 || strings.TrimSpace(sid) == "" {
return AuthIdentity{}, ErrLoginSessionInvalid
}
session, err := model.GetUserSessionCached(sid)
if err != nil {
return AuthIdentity{}, err
}
identity := AuthIdentity{
UserID: userID,
SessionID: sid,
UserAuthVersion: session.UserAuthVersion,
SessionVersion: session.Version,
}
if _, _, err := ValidateLoginSession(identity); err != nil {
return AuthIdentity{}, err
}
return identity, nil
}
// AdvanceCurrentSessionSecurity increments the user's global auth version,
// preserves only the current browser session at a new session version and
// returns a replacement access token. Call after a successful 2FA/passkey
// security-setting mutation that did not already advance AuthVersion.
func AdvanceCurrentSessionSecurity(identity AuthIdentity, reason string) (*AuthBundle, error) {
nextUserAuthVersion, err := model.BumpUserAuthVersion(identity.UserID)
if err != nil {
return nil, err
}
return advanceCurrentSessionToVersion(identity, nextUserAuthVersion, reason)
}
// AdvanceCurrentSessionToUserVersion is used when the security mutation and
// AuthVersion increment were committed in the same transaction (for example,
// a password change).
func AdvanceCurrentSessionToUserVersion(identity AuthIdentity, reason string) (*AuthBundle, error) {
user, err := model.GetUserCache(identity.UserID)
if err != nil {
return nil, err
}
if user.Status != common.UserStatusEnabled || user.AuthVersion <= identity.UserAuthVersion {
return nil, ErrLoginSessionRevoked
}
return advanceCurrentSessionToVersion(identity, user.AuthVersion, reason)
}
func advanceCurrentSessionToVersion(identity AuthIdentity, nextUserAuthVersion int64, reason string) (*AuthBundle, error) {
session, err := model.AdvanceUserSessionAuthVersion(
identity.UserID,
identity.SessionID,
identity.SessionVersion,
identity.UserAuthVersion,
nextUserAuthVersion,
)
if err != nil {
return nil, err
}
if _, err := model.RevokeOtherUserSessions(identity.UserID, identity.SessionID, reason); err != nil {
return nil, err
}
return issueAuthBundle(session, "", true)
}
func RefreshLoginSession(rawRefreshToken, expectedSID, ip, userAgent string) (*AuthBundle, *model.User, error) {
sid, secret, ok := splitRefreshToken(rawRefreshToken)
if !ok {
return nil, nil, ErrRefreshTokenInvalid
}
if expectedSID = strings.TrimSpace(expectedSID); expectedSID != "" && expectedSID != sid {
return nil, nil, ErrLoginSessionMismatch
}
session, err := model.GetUserSessionCached(sid)
if err != nil {
if errors.Is(err, model.ErrUserSessionInactive) {
return nil, nil, ErrLoginSessionRevoked
}
return nil, nil, ErrRefreshTokenInvalid
}
if session.Status != model.UserSessionStatusActive || session.RevokedAt != 0 || session.ExpiresAt <= time.Now().Unix() {
return nil, nil, ErrLoginSessionRevoked
}
userCache, err := model.GetUserCache(session.UserID)
if err != nil {
return nil, nil, err
}
currentUser, err := model.GetUserById(session.UserID, false)
if err != nil {
return nil, nil, err
}
if userCache.Status != common.UserStatusEnabled || userCache.AuthVersion != session.UserAuthVersion ||
currentUser.Status != common.UserStatusEnabled || currentUser.AuthVersion != session.UserAuthVersion {
_, _ = model.RevokeUserSession(session.UserID, session.SID, "user_security_changed")
return nil, nil, ErrLoginSessionRevoked
}
nextSecret := deriveNextRefreshSecret(sid, secret)
rotated, err := model.RotateUserSessionRefresh(session.UserID, sid, hashRefreshSecret(secret), hashRefreshSecret(nextSecret), time.Now().Unix(), RefreshReplayWindow)
if err != nil {
if errors.Is(err, model.ErrUserSessionRefreshRace) && rotated != nil &&
hashRefreshSecret(nextSecret) == rotated.RefreshHash {
bundle, issueErr := issueAuthBundle(rotated, sid+"."+nextSecret, true)
if issueErr != nil {
return nil, nil, issueErr
}
return bundle, currentUser, nil
}
if errors.Is(err, model.ErrUserSessionRefreshReuse) {
return nil, nil, ErrLoginSessionRevoked
}
if errors.Is(err, model.ErrUserSessionRefreshInvalid) {
return nil, nil, ErrRefreshTokenInvalid
}
if errors.Is(err, model.ErrUserSessionRefreshRace) {
return nil, nil, ErrRefreshRace
}
return nil, nil, err
}
rotated.IP = truncateAuthMetadata(ip, 64)
rotated.UserAgent = truncateAuthMetadata(userAgent, 512)
bundle, err := issueAuthBundle(rotated, sid+"."+nextSecret, true)
if err != nil {
return nil, nil, err
}
return bundle, currentUser, nil
}
func RevokeByRefreshToken(rawRefreshToken, expectedSID, reason string) error {
sid, secret, ok := splitRefreshToken(rawRefreshToken)
if !ok {
return nil
}
if expectedSID = strings.TrimSpace(expectedSID); expectedSID != "" && expectedSID != sid {
return ErrLoginSessionMismatch
}
_, err := model.RevokeUserSessionByRefreshHash(sid, hashRefreshSecret(secret), reason)
return err
}
func RefreshTokenSID(rawRefreshToken string) (string, bool) {
sid, _, ok := splitRefreshToken(rawRefreshToken)
return sid, ok
}
func ListLoginSessions(userID int, currentSID string) ([]LoginSessionView, error) {
sessions, err := model.ListActiveUserSessions(userID, currentSID, time.Now().Unix())
if err != nil {
return nil, err
}
views := make([]LoginSessionView, 0, len(sessions))
for i := range sessions {
views = append(views, sessionView(&sessions[i], sessions[i].SID == currentSID))
}
return views, nil
}
func WriteRefreshCookie(c *gin.Context, rawToken string) {
expiresAt := time.Now().Add(LoginSessionTTL)
if sid, _, ok := splitRefreshToken(rawToken); ok {
if session, err := model.GetUserSessionCached(sid); err == nil && session.ExpiresAt > time.Now().Unix() {
expiresAt = time.Unix(session.ExpiresAt, 0)
}
}
maxAge := int(time.Until(expiresAt) / time.Second)
if maxAge < 1 {
maxAge = 1
}
http.SetCookie(c.Writer, &http.Cookie{
Name: RefreshCookieName,
Value: rawToken,
Path: "/api/user/auth",
MaxAge: maxAge,
Expires: expiresAt,
HttpOnly: true,
Secure: common.SessionCookieSecure,
SameSite: http.SameSiteStrictMode,
})
}
func ClearRefreshCookie(c *gin.Context) {
http.SetCookie(c.Writer, &http.Cookie{
Name: RefreshCookieName,
Value: "",
Path: "/api/user/auth",
MaxAge: -1,
Expires: time.Unix(1, 0),
HttpOnly: true,
Secure: common.SessionCookieSecure,
SameSite: http.SameSiteStrictMode,
})
}
func issueAuthBundle(session *model.UserSession, rawRefreshToken string, current bool) (*AuthBundle, error) {
identity := AuthIdentity{
UserID: session.UserID,
SessionID: session.SID,
UserAuthVersion: session.UserAuthVersion,
SessionVersion: session.Version,
}
accessToken, accessExpiresAt, err := IssueAccessToken(identity)
if err != nil {
return nil, err
}
return &AuthBundle{
AccessToken: accessToken,
TokenType: "Bearer",
AccessExpiresAt: accessExpiresAt,
Session: sessionView(session, current),
RefreshToken: rawRefreshToken,
}, nil
}
func sessionView(session *model.UserSession, current bool) LoginSessionView {
return LoginSessionView{
SID: session.SID,
Current: current,
LoginMethod: session.LoginMethod,
IP: session.IP,
UserAgent: session.UserAgent,
CreatedAt: session.CreatedAt,
LastActiveAt: session.LastActiveAt,
ExpiresAt: session.ExpiresAt,
}
}
func splitRefreshToken(raw string) (string, string, bool) {
sid, secret, ok := strings.Cut(strings.TrimSpace(raw), ".")
if !ok || sid == "" || secret == "" || strings.Contains(secret, ".") {
return "", "", false
}
if _, err := uuid.Parse(sid); err != nil {
return "", "", false
}
return sid, secret, true
}
func hashRefreshSecret(secret string) string {
return common.GenerateHMACWithKey(authSigningKey("refresh"), secret)
}
func deriveNextRefreshSecret(sid, currentSecret string) string {
return common.GenerateHMACWithKey(authSigningKey("refresh-rotate"), sid+"."+currentSecret)
}
func truncateAuthMetadata(value string, max int) string {
value = strings.TrimSpace(value)
if len(value) <= max {
return value
}
return value[:max]
}
func authSessionErrorCode(err error) (int, string) {
switch {
case errors.Is(err, model.ErrUserSessionLimit):
return http.StatusConflict, "AUTH_SESSION_LIMIT"
case errors.Is(err, model.ErrUserSessionIssuanceLimit):
return http.StatusTooManyRequests, "AUTH_SESSION_ISSUANCE_LIMIT"
case errors.Is(err, ErrLoginSessionMismatch):
return http.StatusConflict, "AUTH_SESSION_MISMATCH"
case errors.Is(err, ErrRefreshRace):
return http.StatusConflict, "AUTH_REFRESH_RACE"
case errors.Is(err, ErrAuthTokenExpired):
return http.StatusUnauthorized, "AUTH_TOKEN_EXPIRED"
case errors.Is(err, ErrLoginSessionRevoked):
return http.StatusUnauthorized, "AUTH_SESSION_REVOKED"
case errors.Is(err, ErrRefreshTokenInvalid), errors.Is(err, ErrAuthTokenInvalid):
return http.StatusUnauthorized, "AUTH_UNAUTHORIZED"
default:
return http.StatusInternalServerError, "AUTH_INTERNAL_ERROR"
}
}
func AuthSessionErrorCode(err error) (int, string) {
return authSessionErrorCode(err)
}
func FormatAuthError(err error) string {
if err == nil {
return ""
}
return fmt.Sprintf("authentication failed: %v", err)
}
+431
View File
@@ -0,0 +1,431 @@
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")
}
+216
View File
@@ -0,0 +1,216 @@
package service
import (
"crypto/hmac"
"crypto/sha256"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/QuantumNous/new-api/common"
"github.com/golang-jwt/jwt/v5"
"github.com/google/uuid"
)
const (
AccessTokenTTL = 15 * time.Minute
SecurityProofTTL = 5 * time.Minute
LoginSessionTTL = 30 * 24 * time.Hour
RefreshReplayWindow = 30 * time.Second
accessTokenUse = "access"
securityProofTokenUse = "security_proof"
authTokenIssuer = "new-api"
authTokenAudience = "new-api-dashboard"
)
var (
ErrAuthTokenInvalid = errors.New("authentication token is invalid")
ErrAuthTokenExpired = errors.New("authentication token has expired")
ErrProofScope = errors.New("security proof scope mismatch")
ErrProofMethod = errors.New("security proof method mismatch")
)
// AuthIdentity is the server-validated identity attached to dashboard requests.
// Role, status and group are deliberately loaded from the user cache instead of JWT claims.
type AuthIdentity struct {
UserID int
SessionID string
UserAuthVersion int64
SessionVersion int64
}
type authClaims struct {
TokenUse string `json:"token_use"`
SessionID string `json:"sid"`
UserAuthVersion int64 `json:"uv"`
SessionVersion int64 `json:"sv"`
Method string `json:"method,omitempty"`
Scopes []string `json:"scopes,omitempty"`
jwt.RegisteredClaims
}
func authSigningKey(purpose string) []byte {
mac := hmac.New(sha256.New, []byte(common.SessionSecret))
_, _ = mac.Write([]byte("new-api/auth/" + purpose + "/v1"))
return mac.Sum(nil)
}
func IssueAccessToken(identity AuthIdentity) (string, int64, error) {
if identity.UserID <= 0 || identity.SessionID == "" || identity.UserAuthVersion <= 0 || identity.SessionVersion <= 0 {
return "", 0, ErrAuthTokenInvalid
}
now := time.Now()
expiresAt := now.Add(AccessTokenTTL)
claims := authClaims{
TokenUse: accessTokenUse,
SessionID: identity.SessionID,
UserAuthVersion: identity.UserAuthVersion,
SessionVersion: identity.SessionVersion,
RegisteredClaims: jwt.RegisteredClaims{
Issuer: authTokenIssuer,
Subject: strconv.Itoa(identity.UserID),
Audience: jwt.ClaimStrings{authTokenAudience},
ExpiresAt: jwt.NewNumericDate(expiresAt),
NotBefore: jwt.NewNumericDate(now.Add(-5 * time.Second)),
IssuedAt: jwt.NewNumericDate(now),
ID: uuid.NewString(),
},
}
signed, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(authSigningKey(accessTokenUse))
return signed, expiresAt.Unix(), err
}
func ParseAccessToken(raw string) (AuthIdentity, error) {
claims, err := parseAuthClaims(raw, accessTokenUse, authSigningKey(accessTokenUse))
if err != nil {
return AuthIdentity{}, err
}
userID, err := strconv.Atoi(claims.Subject)
if err != nil || userID <= 0 || claims.SessionID == "" || claims.UserAuthVersion <= 0 || claims.SessionVersion <= 0 {
return AuthIdentity{}, ErrAuthTokenInvalid
}
return AuthIdentity{
UserID: userID,
SessionID: claims.SessionID,
UserAuthVersion: claims.UserAuthVersion,
SessionVersion: claims.SessionVersion,
}, nil
}
// ParseDashboardAccessToken distinguishes new-api dashboard JWTs from opaque
// credentials. A token carrying the dashboard issuer, audience and a known
// token use is always treated as internal, even when its signature, lifetime
// or requested purpose is invalid, so it can never fall through to PAT or
// relay-token authentication.
func ParseDashboardAccessToken(raw string) (identity AuthIdentity, internal bool, err error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return AuthIdentity{}, false, nil
}
claims := &authClaims{}
parsed, _, parseErr := jwt.NewParser().ParseUnverified(raw, claims)
if parseErr != nil || parsed == nil {
return AuthIdentity{}, false, nil
}
audienceMatches := false
for _, audience := range claims.Audience {
if audience == authTokenAudience {
audienceMatches = true
break
}
}
knownTokenUse := claims.TokenUse == accessTokenUse || claims.TokenUse == securityProofTokenUse
if claims.Issuer != authTokenIssuer || !audienceMatches || !knownTokenUse {
return AuthIdentity{}, false, nil
}
identity, err = ParseAccessToken(raw)
return identity, true, err
}
func IssueSecurityProof(identity AuthIdentity, method string, scopes []string) (string, int64, error) {
method = strings.TrimSpace(method)
if identity.UserID <= 0 || identity.SessionID == "" || identity.UserAuthVersion <= 0 || identity.SessionVersion <= 0 || method == "" || len(scopes) == 0 {
return "", 0, ErrAuthTokenInvalid
}
now := time.Now()
expiresAt := now.Add(SecurityProofTTL)
claims := authClaims{
TokenUse: securityProofTokenUse,
SessionID: identity.SessionID,
UserAuthVersion: identity.UserAuthVersion,
SessionVersion: identity.SessionVersion,
Method: method,
Scopes: append([]string(nil), scopes...),
RegisteredClaims: jwt.RegisteredClaims{
Issuer: authTokenIssuer,
Subject: strconv.Itoa(identity.UserID),
Audience: jwt.ClaimStrings{authTokenAudience},
ExpiresAt: jwt.NewNumericDate(expiresAt),
NotBefore: jwt.NewNumericDate(now.Add(-5 * time.Second)),
IssuedAt: jwt.NewNumericDate(now),
ID: uuid.NewString(),
},
}
signed, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(authSigningKey(securityProofTokenUse))
return signed, expiresAt.Unix(), err
}
func VerifySecurityProof(raw string, identity AuthIdentity, requiredScope string, allowedMethods []string) (string, error) {
claims, err := parseAuthClaims(raw, securityProofTokenUse, authSigningKey(securityProofTokenUse))
if err != nil {
return "", err
}
userID, err := strconv.Atoi(claims.Subject)
if err != nil || userID != identity.UserID || claims.SessionID != identity.SessionID || claims.UserAuthVersion != identity.UserAuthVersion || claims.SessionVersion != identity.SessionVersion {
return "", ErrAuthTokenInvalid
}
methodAllowed := len(allowedMethods) == 0
for _, method := range allowedMethods {
if hmac.Equal([]byte(claims.Method), []byte(method)) {
methodAllowed = true
break
}
}
if !methodAllowed {
return "", ErrProofMethod
}
if requiredScope != "" {
found := false
for _, scope := range claims.Scopes {
if hmac.Equal([]byte(scope), []byte(requiredScope)) {
found = true
break
}
}
if !found {
return "", ErrProofScope
}
}
return claims.Method, nil
}
func parseAuthClaims(raw, expectedUse string, key []byte) (*authClaims, error) {
raw = strings.TrimSpace(raw)
if raw == "" {
return nil, ErrAuthTokenInvalid
}
claims := &authClaims{}
parsed, err := jwt.ParseWithClaims(raw, claims, func(token *jwt.Token) (any, error) {
if token.Method.Alg() != jwt.SigningMethodHS256.Alg() {
return nil, fmt.Errorf("%w: unexpected signing method", ErrAuthTokenInvalid)
}
return key, nil
}, jwt.WithValidMethods([]string{jwt.SigningMethodHS256.Alg()}), jwt.WithIssuer(authTokenIssuer), jwt.WithAudience(authTokenAudience), jwt.WithExpirationRequired(), jwt.WithIssuedAt(), jwt.WithLeeway(5*time.Second))
if err != nil {
if errors.Is(err, jwt.ErrTokenExpired) {
return nil, ErrAuthTokenExpired
}
return nil, fmt.Errorf("%w: %v", ErrAuthTokenInvalid, err)
}
if !parsed.Valid || claims.TokenUse != expectedUse || claims.ID == "" || claims.IssuedAt == nil || claims.NotBefore == nil {
return nil, ErrAuthTokenInvalid
}
return claims, nil
}
+140
View File
@@ -0,0 +1,140 @@
package service
import (
"errors"
"testing"
"time"
"github.com/QuantumNous/new-api/common"
"github.com/golang-jwt/jwt/v5"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func useTestSessionSecret(t *testing.T) {
t.Helper()
previous := common.SessionSecret
common.SessionSecret = "test-session-secret-with-sufficient-entropy"
t.Cleanup(func() { common.SessionSecret = previous })
}
func TestAccessTokenRoundTripAndPurposeIsolation(t *testing.T) {
useTestSessionSecret(t)
identity := AuthIdentity{UserID: 42, SessionID: "session-1", UserAuthVersion: 3, SessionVersion: 2}
token, expiresAt, err := IssueAccessToken(identity)
require.NoError(t, err)
assert.Positive(t, expiresAt)
parsed, err := ParseAccessToken(token)
require.NoError(t, err)
assert.Equal(t, identity, parsed)
proof, _, err := IssueSecurityProof(identity, "2fa", []string{"channel.key.read"})
require.NoError(t, err)
_, err = ParseAccessToken(proof)
assert.ErrorIs(t, err, ErrAuthTokenInvalid)
}
func TestAccessTokenRejectsTampering(t *testing.T) {
useTestSessionSecret(t)
identity := AuthIdentity{UserID: 42, SessionID: "session-1", UserAuthVersion: 1, SessionVersion: 1}
token, _, err := IssueAccessToken(identity)
require.NoError(t, err)
tamperAt := len(token) - 2
replacement := "x"
if token[tamperAt] == 'x' {
replacement = "y"
}
tampered := token[:tamperAt] + replacement + token[tamperAt+1:]
_, err = ParseAccessToken(tampered)
assert.ErrorIs(t, err, ErrAuthTokenInvalid)
_, internal, err := ParseDashboardAccessToken(tampered)
assert.True(t, internal)
assert.ErrorIs(t, err, ErrAuthTokenInvalid)
}
func TestDashboardAccessTokenClassification(t *testing.T) {
useTestSessionSecret(t)
identity, internal, err := ParseDashboardAccessToken("opaque.key.with-dots")
require.NoError(t, err)
assert.False(t, internal)
assert.Empty(t, identity)
external := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
"iss": "external-issuer",
"aud": authTokenAudience,
"exp": time.Now().Add(time.Minute).Unix(),
})
externalRaw, err := external.SignedString([]byte("external-secret"))
require.NoError(t, err)
_, internal, err = ParseDashboardAccessToken(externalRaw)
require.NoError(t, err)
assert.False(t, internal)
unknownUse := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
"iss": authTokenIssuer,
"aud": authTokenAudience,
"token_use": "third_party",
"exp": time.Now().Add(time.Minute).Unix(),
})
unknownUseRaw, err := unknownUse.SignedString([]byte("external-secret"))
require.NoError(t, err)
_, internal, err = ParseDashboardAccessToken(unknownUseRaw)
require.NoError(t, err)
assert.False(t, internal)
proof, _, err := IssueSecurityProof(AuthIdentity{
UserID: 42, SessionID: "session-1", UserAuthVersion: 1, SessionVersion: 1,
}, "2fa", []string{"channel.key.read"})
require.NoError(t, err)
_, internal, err = ParseDashboardAccessToken(proof)
assert.True(t, internal)
assert.ErrorIs(t, err, ErrAuthTokenInvalid)
expiredClaims := authClaims{
TokenUse: accessTokenUse,
SessionID: "expired-session",
UserAuthVersion: 1,
SessionVersion: 1,
RegisteredClaims: jwt.RegisteredClaims{
Issuer: authTokenIssuer,
Subject: "42",
Audience: jwt.ClaimStrings{authTokenAudience},
ExpiresAt: jwt.NewNumericDate(time.Now().Add(-time.Minute)),
NotBefore: jwt.NewNumericDate(time.Now().Add(-2 * time.Minute)),
IssuedAt: jwt.NewNumericDate(time.Now().Add(-2 * time.Minute)),
ID: "expired-token",
},
}
expired, err := jwt.NewWithClaims(jwt.SigningMethodHS256, expiredClaims).SignedString(authSigningKey(accessTokenUse))
require.NoError(t, err)
_, internal, err = ParseDashboardAccessToken(expired)
assert.True(t, internal)
assert.ErrorIs(t, err, ErrAuthTokenExpired)
}
func TestSecurityProofBindsIdentityMethodAndScope(t *testing.T) {
useTestSessionSecret(t)
identity := AuthIdentity{UserID: 42, SessionID: "session-1", UserAuthVersion: 3, SessionVersion: 2}
proof, _, err := IssueSecurityProof(identity, "2fa", []string{"channel.key.read"})
require.NoError(t, err)
method, err := VerifySecurityProof(proof, identity, "channel.key.read", []string{"2fa", "passkey"})
require.NoError(t, err)
assert.Equal(t, "2fa", method)
_, err = VerifySecurityProof(proof, identity, "passkey.delete", []string{"2fa"})
assert.ErrorIs(t, err, ErrProofScope)
_, err = VerifySecurityProof(proof, identity, "channel.key.read", []string{"passkey"})
assert.ErrorIs(t, err, ErrProofMethod)
otherSession := identity
otherSession.SessionID = "session-2"
_, err = VerifySecurityProof(proof, otherSession, "channel.key.read", []string{"2fa"})
assert.True(t, errors.Is(err, ErrAuthTokenInvalid))
}
-6
View File
@@ -16,12 +16,6 @@ import (
webauthn "github.com/go-webauthn/webauthn/webauthn"
)
const (
RegistrationSessionKey = "passkey_registration_session"
LoginSessionKey = "passkey_login_session"
VerifySessionKey = "passkey_verify_session"
)
// BuildWebAuthn constructs a WebAuthn instance using the current passkey settings and request context.
func BuildWebAuthn(r *http.Request) (*webauthn.WebAuthn, error) {
settings := system_setting.GetPasskeySettings()
+46 -35
View File
@@ -1,50 +1,61 @@
package passkey
import (
"encoding/json"
"errors"
"time"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
webauthn "github.com/go-webauthn/webauthn/webauthn"
)
var errSessionNotFound = errors.New("Passkey 会话不存在或已过期")
func SaveSessionData(c *gin.Context, key string, data *webauthn.SessionData) error {
session := sessions.Default(c)
if data == nil {
session.Delete(key)
return session.Save()
}
payload, err := json.Marshal(data)
if err != nil {
return err
}
session.Set(key, string(payload))
return session.Save()
const passkeyFlowTTL = 5 * time.Minute
type flowPayload struct {
SessionData webauthn.SessionData `json:"session_data"`
Scope string `json:"scope,omitempty"`
}
func PopSessionData(c *gin.Context, key string) (*webauthn.SessionData, error) {
session := sessions.Default(c)
raw := session.Get(key)
if raw == nil {
return nil, errSessionNotFound
func CreateSessionDataFlow(purpose string, userID int, sessionID, scope string, data *webauthn.SessionData) (string, int64, error) {
if data == nil {
return "", 0, errors.New("Passkey 会话数据不能为空")
}
session.Delete(key)
_ = session.Save()
var data webauthn.SessionData
switch value := raw.(type) {
case string:
if err := json.Unmarshal([]byte(value), &data); err != nil {
return nil, err
}
case []byte:
if err := json.Unmarshal(value, &data); err != nil {
return nil, err
}
default:
return nil, errors.New("Passkey 会话格式无效")
payload, err := common.Marshal(flowPayload{SessionData: *data, Scope: scope})
if err != nil {
return "", 0, err
}
return &data, nil
expiresAt := time.Now().Add(passkeyFlowTTL)
token, _, err := model.CreateAuthFlow(model.AuthFlowCreate{
Purpose: purpose,
UserId: userID,
SessionId: sessionID,
Payload: string(payload),
ExpiresAt: expiresAt,
})
if err != nil {
return "", 0, err
}
return token, expiresAt.Unix(), nil
}
func PopSessionDataFlow(token, purpose string, userID int, sessionID string) (*webauthn.SessionData, string, error) {
flow, err := model.ConsumeAuthFlow(token, model.AuthFlowMatch{
Purpose: purpose,
UserId: userID,
SessionId: sessionID,
})
if err != nil {
if errors.Is(err, model.ErrAuthFlowInvalid) || errors.Is(err, model.ErrAuthFlowExpired) || errors.Is(err, model.ErrAuthFlowConsumed) {
return nil, "", errSessionNotFound
}
return nil, "", err
}
var payload flowPayload
if err := common.UnmarshalJsonStr(flow.Payload, &payload); err != nil {
return nil, "", err
}
return &payload.SessionData, payload.Scope, nil
}
+2 -2
View File
@@ -470,7 +470,7 @@ func checkAndSendQuotaNotify(relayInfo *relaycommon.RelayInfo, quota int, preCon
}
if quotaTooLow {
prompt := "您的额度即将用尽"
topUpLink := PaymentReturnURL("/console/topup")
topUpLink := PaymentReturnURL("/wallet")
// 根据通知方式生成不同的内容格式
var content string
@@ -524,7 +524,7 @@ func checkAndSendSubscriptionQuotaNotify(relayInfo *relaycommon.RelayInfo) {
}
prompt := "您的订阅额度即将用尽"
topUpLink := PaymentReturnURL("/console/topup")
topUpLink := PaymentReturnURL("/wallet")
var content string
var values []interface{}
+1 -2
View File
@@ -3,11 +3,10 @@ package service
import (
"strings"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/setting/system_setting"
)
func PaymentReturnURL(suffix string) string {
base := strings.TrimRight(system_setting.ServerAddress, "/")
return base + common.ThemeAwarePath(suffix)
return base + suffix
}
+16
View File
@@ -0,0 +1,16 @@
package service
import (
"testing"
"github.com/QuantumNous/new-api/setting/system_setting"
"github.com/stretchr/testify/assert"
)
func TestPaymentReturnURLUsesSuppliedDefaultDashboardPath(t *testing.T) {
previousAddress := system_setting.ServerAddress
system_setting.ServerAddress = "https://dashboard.example.com/"
t.Cleanup(func() { system_setting.ServerAddress = previousAddress })
assert.Equal(t, "https://dashboard.example.com/wallet", PaymentReturnURL("/wallet"))
}