Files
jiu/backend/internal/service/session_test.go
wangjia 15ef71734d
Deploy Server / release-deploy-server (push) Successful in 55s
chore: release server-v1.0.64
Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-20 10:13:49 +08:00

178 lines
6.5 KiB
Go

package service
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/wangjia/jiu/backend/config"
"github.com/wangjia/jiu/backend/internal/model"
"github.com/wangjia/jiu/backend/testutil"
)
// 达到总配额时拒绝新登录(不踢最旧),已有会话保持在线。
func TestLogin_TotalQuota_RejectsWhenFull(t *testing.T) {
db := testutil.SetupTestDB()
old := config.C.Session.LimitTotal
config.C.Session.LimitTotal = 2
defer func() { config.C.Session.LimitTotal = old }()
shop := testutil.CreateTestShop(db, "SESS01")
testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
dev := DeviceInfo{Platform: "windows"}
p1, _, err := svc.Login("SESS01", "admin", "password123", dev)
require.NoError(t, err)
p2, _, err := svc.Login("SESS01", "admin", "password123", dev)
require.NoError(t, err)
// 总配额 2 已满,第 3 次登录应被拒绝
_, _, err = svc.Login("SESS01", "admin", "password123", dev)
assert.ErrorIs(t, err, ErrDeviceLimitReached)
// 仍为 2 个未撤销会话,两者都未被踢
var active int64
db.Model(&model.UserSession{}).Where("shop_id = ? AND revoked_at IS NULL", shop.ID).Count(&active)
assert.Equal(t, int64(2), active)
_, err = svc.RefreshTokens(p1.RefreshToken)
assert.NoError(t, err)
_, err = svc.RefreshTokens(p2.RefreshToken)
assert.NoError(t, err)
}
// 总配额跨全部平台合并计数:不同平台类共用一个总上限,满额即拒。
func TestLogin_TotalQuota_CountsAcrossClasses(t *testing.T) {
db := testutil.SetupTestDB()
old := config.C.Session.LimitTotal
config.C.Session.LimitTotal = 2
defer func() { config.C.Session.LimitTotal = old }()
shop := testutil.CreateTestShop(db, "SESS02")
testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
// 桌面 + 移动 占满总配额 2
_, _, err := svc.Login("SESS02", "admin", "password123", DeviceInfo{Platform: "windows"})
require.NoError(t, err)
_, _, err = svc.Login("SESS02", "admin", "password123", DeviceInfo{Platform: "android"})
require.NoError(t, err)
// 第三个不同平台(web)应被总配额拒绝
_, _, err = svc.Login("SESS02", "admin", "password123", DeviceInfo{Platform: "web"})
assert.ErrorIs(t, err, ErrDeviceLimitReached)
var active int64
db.Model(&model.UserSession{}).Where("shop_id = ? AND revoked_at IS NULL", shop.ID).Count(&active)
assert.Equal(t, int64(2), active) // 仅前两个在线
}
// 配额为 0 的平台拒绝登录(如禁 web)。
func TestLogin_PlatformDisabled(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "SESS03")
testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
old := config.C.Session.LimitWeb
config.C.Session.LimitWeb = 0
defer func() { config.C.Session.LimitWeb = old }()
_, _, err := svc.Login("SESS03", "admin", "password123", DeviceInfo{Platform: "web"})
assert.ErrorIs(t, err, ErrPlatformNotAllowed)
}
// 每店 session_policy["total"] 覆盖全局总配额默认。
func TestLogin_PerShopPolicyOverride(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "SESS04")
// 该店总配额限 1
db.Model(&model.Shop{}).Where("id = ?", shop.ID).
Update("custom_fields", model.JSON{"session_policy": map[string]interface{}{"total": float64(1)}})
testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
p1, _, err := svc.Login("SESS04", "admin", "password123", DeviceInfo{Platform: "windows"})
require.NoError(t, err)
// 总配额 1 已满,第 2 次登录(不同平台)应被拒绝,p1 保持在线
_, _, err = svc.Login("SESS04", "admin", "password123", DeviceInfo{Platform: "android"})
assert.ErrorIs(t, err, ErrDeviceLimitReached)
_, err = svc.RefreshTokens(p1.RefreshToken)
assert.NoError(t, err)
}
// 被禁用用户无法用 refresh token 续命(修复历史漏洞)。
func TestRefresh_DisabledUserRejected(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "SESS05")
user := testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
pair, _, err := svc.Login("SESS05", "admin", "password123", DeviceInfo{Platform: "windows"})
require.NoError(t, err)
db.Model(user).Update("is_active", false)
_, err = svc.RefreshTokens(pair.RefreshToken)
assert.ErrorIs(t, err, ErrUserInactive)
}
// 管理员强制下线后,该会话 token 无法续期。
func TestForceLogout_RevokesSession(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "SESS06")
testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
pair, _, err := svc.Login("SESS06", "admin", "password123", DeviceInfo{Platform: "windows"})
require.NoError(t, err)
views, err := svc.ListSessions(shop.ID, "")
require.NoError(t, err)
require.Len(t, views, 1)
assert.True(t, views[0].Online)
require.NoError(t, svc.ForceLogout(shop.ID, views[0].ID, 0))
_, err = svc.RefreshTokens(pair.RefreshToken)
assert.ErrorIs(t, err, ErrSessionRevoked)
// 列表中不再出现
views, err = svc.ListSessions(shop.ID, "")
require.NoError(t, err)
assert.Len(t, views, 0)
}
// 强制下线跨店隔离:不能下线别店会话。
func TestForceLogout_TenantIsolation(t *testing.T) {
db := testutil.SetupTestDB()
shopA := testutil.CreateTestShop(db, "SESS07A")
shopB := testutil.CreateTestShop(db, "SESS07B")
testutil.CreateTestUser(db, shopA.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
_, _, err := svc.Login("SESS07A", "admin", "password123", DeviceInfo{Platform: "windows"})
require.NoError(t, err)
views, _ := svc.ListSessions(shopA.ID, "")
require.Len(t, views, 1)
// 用 shopB 的 shopID 尝试下线 shopA 的会话 → 找不到
err = svc.ForceLogout(shopB.ID, views[0].ID, 0)
assert.Error(t, err)
}
// 连续登录失败达到阈值后锁定。
func TestLogin_LockoutAfterFailures(t *testing.T) {
db := testutil.SetupTestDB()
shop := testutil.CreateTestShop(db, "SESS08")
testutil.CreateTestUser(db, shop.ID, "admin", "password123", "admin")
svc := NewAuthService(db)
for i := 0; i < config.C.Session.MaxFailures; i++ {
_, _, err := svc.Login("SESS08", "admin", "wrong", DeviceInfo{Platform: "windows"})
assert.ErrorIs(t, err, ErrInvalidCredentials)
}
// 锁定后即便密码正确也被拒
_, _, err := svc.Login("SESS08", "admin", "password123", DeviceInfo{Platform: "windows"})
assert.ErrorIs(t, err, ErrTooManyAttempts)
}