247 lines
6.6 KiB
Go
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)
|
|
}
|
|
}
|