Files
usernode/internal/service/user_test.go
T

247 lines
6.6 KiB
Go

package service
import (
"context"
"errors"
"sync"
"testing"
"time"
"gorm.io/gorm"
"ws_usernode/internal/config"
"ws_usernode/internal/model"
"ws_usernode/internal/system"
)
// fakeSys 内存版 system.Manager:记录已创建的账号,便于断言与真实系统隔离。
type fakeSys struct {
mu sync.Mutex
accounts map[string]bool
keys map[string][]system.Key // username -> 最后一次同步的密钥
lastCmd string
syncErr error // 注入 SyncAuthorizedKeys 失败(测试回滚)
}
func newFakeSys() *fakeSys {
return &fakeSys{accounts: map[string]bool{}, keys: map[string][]system.Key{}}
}
func (f *fakeSys) CreateUser(_ context.Context, acc system.Account) error {
f.mu.Lock()
defer f.mu.Unlock()
f.accounts[acc.Username] = true
f.lastCmd = "useradd " + acc.Username
return nil
}
func (f *fakeSys) RemoveUser(_ context.Context, username string) error {
f.mu.Lock()
defer f.mu.Unlock()
delete(f.accounts, username)
f.lastCmd = "userdel " + username
return nil
}
func (f *fakeSys) SetLock(_ context.Context, username string, locked bool) error {
f.mu.Lock()
defer f.mu.Unlock()
f.lastCmd = "passwd " + username
return nil
}
func (f *fakeSys) Exists(_ context.Context, username string) (bool, error) {
f.mu.Lock()
defer f.mu.Unlock()
return f.accounts[username], nil
}
func (f *fakeSys) SyncAuthorizedKeys(_ context.Context, username string, keys []system.Key) error {
f.mu.Lock()
defer f.mu.Unlock()
if f.syncErr != nil {
return f.syncErr
}
f.keys[username] = keys
f.lastCmd = "sync-keys " + username
return nil
}
func (f *fakeSys) has(username string) bool {
f.mu.Lock()
defer f.mu.Unlock()
return f.accounts[username]
}
func testDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := model.Open("sqlite", ":memory:", false)
if err != nil {
t.Fatalf("open test db: %v", err)
}
if err := model.Migrate(db); err != nil {
t.Fatalf("migrate: %v", err)
}
return db
}
func testConfig() *config.Config {
cfg := config.Default()
cfg.System.DryRun = false // 测试直接走 fakeSys,不依赖 dry-run
return cfg
}
func mustAdmin(t *testing.T, svc *AdminService) *model.AdminUser {
t.Helper()
adm, err := svc.Create(context.Background(), "root", "Passw0rd", "root@example.com")
if err != nil {
t.Fatalf("create admin: %v", err)
}
return adm
}
func TestAdminServiceLogin(t *testing.T) {
db := testDB(t)
svc := NewAdminService(db)
adm := mustAdmin(t, svc)
got, err := svc.Login(context.Background(), adm.Username, "Passw0rd")
if err != nil {
t.Fatalf("login: %v", err)
}
if got.ID != adm.ID {
t.Fatalf("login returned wrong admin")
}
if _, err := svc.Login(context.Background(), adm.Username, "wrong"); !errors.Is(err, ErrBadCredentials) {
t.Fatalf("bad password err = %v, want ErrBadCredentials", err)
}
if _, err := svc.Login(context.Background(), "ghost", "Passw0rd"); !errors.Is(err, ErrBadCredentials) {
t.Fatalf("missing user err = %v, want ErrBadCredentials", err)
}
}
func TestUserServiceLifecycle(t *testing.T) {
db := testDB(t)
sys := newFakeSys()
svc := NewUserService(db, sys, testConfig())
ctx := context.Background()
u, err := svc.Create(ctx, "zhangsan", "zs@example.com", "prof.li", "科研", 90*24*time.Hour, 1)
if err != nil {
t.Fatalf("create: %v", err)
}
if u.Username != "ext_zhangsan" {
t.Fatalf("username = %q, want ext_zhangsan", u.Username)
}
if !sys.has("ext_zhangsan") {
t.Fatal("system account should exist after create")
}
if u.ExpireAt == nil {
t.Fatal("expire_at should be set")
}
// 重复创建冲突
if _, err := svc.Create(ctx, "zhangsan", "x@example.com", "", "", 0, 1); !errors.Is(err, ErrUserExists) {
t.Fatalf("duplicate create err = %v, want ErrUserExists", err)
}
// 列表
users, total, err := svc.List(ctx, UserFilter{Page: 1, PageSize: 20})
if err != nil {
t.Fatalf("list: %v", err)
}
if total != 1 || len(users) != 1 {
t.Fatalf("list total=%d len=%d, want 1/1", total, len(users))
}
// 更新
email := "new@example.com"
supp := "prof.wang"
u2, err := svc.Update(ctx, u.ID, &email, &supp, nil)
if err != nil {
t.Fatalf("update: %v", err)
}
if u2.Email != "new@example.com" || u2.Supervisor != "prof.wang" {
t.Fatalf("update not applied: %+v", u2)
}
// 禁用 → 系统密钥清空
if err := svc.Disable(ctx, u.ID); err != nil {
t.Fatalf("disable: %v", err)
}
if u2, _ := svc.GetByID(ctx, u.ID); u2.Status != model.UserStatusDisabled {
t.Fatalf("status = %q, want disabled", u2.Status)
}
// 启用 → 恢复 active
if err := svc.Enable(ctx, u.ID); err != nil {
t.Fatalf("enable: %v", err)
}
if u2, _ := svc.GetByID(ctx, u.ID); u2.Status != model.UserStatusActive {
t.Fatalf("status = %q, want active", u2.Status)
}
// 延期
if _, err := svc.Extend(ctx, u.ID, 30); err != nil {
t.Fatalf("extend: %v", err)
}
// 删除 → 系统账号移除
if err := svc.Delete(ctx, u.ID); err != nil {
t.Fatalf("delete: %v", err)
}
if sys.has("ext_zhangsan") {
t.Fatal("system account should be removed after delete")
}
if _, err := svc.GetByID(ctx, u.ID); !errors.Is(err, ErrUserNotFound) {
t.Fatalf("get after delete err = %v, want ErrUserNotFound", err)
}
}
func TestUserServiceEnableExpired(t *testing.T) {
db := testDB(t)
sys := newFakeSys()
svc := NewUserService(db, sys, testConfig())
ctx := context.Background()
u, err := svc.Create(ctx, "lisi", "ls@example.com", "", "", 0, 1)
if err != nil {
t.Fatalf("create: %v", err)
}
// 强制置为过期
past := time.Now().Add(-time.Hour)
if err := db.Model(&model.User{}).Where("id = ?", u.ID).Update("expire_at", &past).Error; err != nil {
t.Fatalf("force expire: %v", err)
}
db.Model(&model.User{}).Where("id = ?", u.ID).Update("status", model.UserStatusExpired)
if err := svc.Enable(ctx, u.ID); !errors.Is(err, ErrUserExpired) {
t.Fatalf("enable expired err = %v, want ErrUserExpired", err)
}
// 延期可恢复
if _, err := svc.Extend(ctx, u.ID, 30); err != nil {
t.Fatalf("extend expired: %v", err)
}
if u2, _ := svc.GetByID(ctx, u.ID); u2.Status != model.UserStatusActive {
t.Fatalf("status after extend = %q, want active", u2.Status)
}
}
func TestUserServiceSystemAccountMissing(t *testing.T) {
db := testDB(t)
sys := newFakeSys()
cfg := testConfig()
cfg.System.DryRun = false
svc := NewUserService(db, sys, cfg)
ctx := context.Background()
// 手动插一条 DB 记录,但系统账号不存在
u := &model.User{Username: "ext_orphan", Email: "o@example.com", Status: model.UserStatusActive, Shell: "/bin/sh"}
if err := db.Create(u).Error; err != nil {
t.Fatalf("insert: %v", err)
}
if err := svc.Disable(ctx, u.ID); !errors.Is(err, ErrSystemAccountMissing) {
t.Fatalf("disable missing account err = %v, want ErrSystemAccountMissing", err)
}
}