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:
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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{}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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"))
|
||||
}
|
||||
Reference in New Issue
Block a user